Skip to content

Commit 1395bae

Browse files
committed
add DeepSeek4 HyperHead, check for DEEPSEEK4 decoder block type in moe.py, and match maxtext sinkhorn implementation to HF reference.
1 parent 409f43b commit 1395bae

9 files changed

Lines changed: 943 additions & 24 deletions

File tree

src/maxtext/checkpoint_conversion/utils/param_mapping.py

Lines changed: 229 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3858,6 +3858,232 @@ def reshape_vision_attn_out(input_tensor, target_shape):
38583858

38593859

38603860
# {maxtext model name: {maxtext weight name: hf weight name}}
3861+
3862+
def DEEPSEEKV4_MAXTEXT_TO_HF_PARAM_MAPPING(config, maxtext_config, scan_layers=False):
3863+
n_layers = config["num_hidden_layers"]
3864+
num_experts = config.get("n_routed_experts", 8)
3865+
3866+
mapping = {
3867+
"params-token_embedder-embedding": "model.embed_tokens.weight",
3868+
"params-decoder-decoder_norm-scale": "model.norm.weight",
3869+
"params-decoder-logits_dense-kernel": "head.weight",
3870+
"params-decoder-hc_head-hc_fn": "model.hc_head.hc_fn",
3871+
"params-decoder-hc_head-hc_base": "model.hc_head.hc_base",
3872+
"params-decoder-hc_head-hc_scale": "model.hc_head.hc_scale",
3873+
}
3874+
3875+
def add_layer_mapping(mt_layer_path, hf_layer_indices):
3876+
is_list = isinstance(hf_layer_indices, list)
3877+
3878+
def get_hf_key(subpath):
3879+
if subpath is None:
3880+
return None
3881+
if is_list:
3882+
return [f"model.layers.{idx}.{subpath}" for idx in hf_layer_indices]
3883+
else:
3884+
return f"model.layers.{hf_layer_indices}.{subpath}"
3885+
3886+
def get_hf_expert_keys(expert_subpath_template):
3887+
if is_list:
3888+
return [
3889+
[f"model.layers.{idx}.mlp.experts.{e}.{expert_subpath_template}" for idx in hf_layer_indices]
3890+
for e in range(num_experts)
3891+
]
3892+
else:
3893+
return [f"model.layers.{hf_layer_indices}.mlp.experts.{e}.{expert_subpath_template}" for e in range(num_experts)]
3894+
3895+
layer_map = {
3896+
f"{mt_layer_path}-pre_self_attention_layer_norm-scale": get_hf_key("input_layernorm.weight"),
3897+
f"{mt_layer_path}-post_self_attention_layer_norm-scale": get_hf_key("post_attention_layernorm.weight"),
3898+
3899+
# Attention
3900+
f"{mt_layer_path}-self_attention-wq_a-kernel": get_hf_key("self_attn.q_a_proj.weight"),
3901+
f"{mt_layer_path}-self_attention-q_norm-scale": get_hf_key("self_attn.q_a_norm.weight"),
3902+
f"{mt_layer_path}-self_attention-wq_b-kernel": get_hf_key("self_attn.q_b_proj.weight"),
3903+
f"{mt_layer_path}-self_attention-wkv-kernel": get_hf_key("self_attn.kv_proj.weight"),
3904+
f"{mt_layer_path}-self_attention-kv_norm-scale": get_hf_key("self_attn.kv_norm.weight"),
3905+
f"{mt_layer_path}-self_attention-sinks": get_hf_key("self_attn.sinks"),
3906+
f"{mt_layer_path}-self_attention-o_a_proj-kernel": get_hf_key("self_attn.o_a_proj.weight"),
3907+
f"{mt_layer_path}-self_attention-o_b_proj-kernel": get_hf_key("self_attn.o_b_proj.weight"),
3908+
3909+
# mHC Attention
3910+
f"{mt_layer_path}-mhc_attention-mhc_norm-scale": None,
3911+
f"{mt_layer_path}-mhc_attention-pre_alpha": get_hf_key("attn_hc.fn"),
3912+
f"{mt_layer_path}-mhc_attention-post_alpha": get_hf_key("attn_hc.fn"),
3913+
f"{mt_layer_path}-mhc_attention-res_alpha": get_hf_key("attn_hc.fn"),
3914+
f"{mt_layer_path}-mhc_attention-pre_beta": get_hf_key("attn_hc.base"),
3915+
f"{mt_layer_path}-mhc_attention-post_beta": get_hf_key("attn_hc.base"),
3916+
f"{mt_layer_path}-mhc_attention-res_beta": get_hf_key("attn_hc.base"),
3917+
f"{mt_layer_path}-mhc_attention-pre_alpha_scale": get_hf_key("attn_hc.scale"),
3918+
f"{mt_layer_path}-mhc_attention-post_alpha_scale": get_hf_key("attn_hc.scale"),
3919+
f"{mt_layer_path}-mhc_attention-res_alpha_scale": get_hf_key("attn_hc.scale"),
3920+
3921+
# mHC MLP
3922+
f"{mt_layer_path}-mhc_mlp-mhc_norm-scale": None,
3923+
f"{mt_layer_path}-mhc_mlp-pre_alpha": get_hf_key("ffn_hc.fn"),
3924+
f"{mt_layer_path}-mhc_mlp-post_alpha": get_hf_key("ffn_hc.fn"),
3925+
f"{mt_layer_path}-mhc_mlp-res_alpha": get_hf_key("ffn_hc.fn"),
3926+
f"{mt_layer_path}-mhc_mlp-pre_beta": get_hf_key("ffn_hc.base"),
3927+
f"{mt_layer_path}-mhc_mlp-post_beta": get_hf_key("ffn_hc.base"),
3928+
f"{mt_layer_path}-mhc_mlp-res_beta": get_hf_key("ffn_hc.base"),
3929+
f"{mt_layer_path}-mhc_mlp-pre_alpha_scale": get_hf_key("ffn_hc.scale"),
3930+
f"{mt_layer_path}-mhc_mlp-post_alpha_scale": get_hf_key("ffn_hc.scale"),
3931+
f"{mt_layer_path}-mhc_mlp-res_alpha_scale": get_hf_key("ffn_hc.scale"),
3932+
3933+
# MoE Block
3934+
f"{mt_layer_path}-mlp-MoeBlock_0-gate-kernel": get_hf_key("mlp.gate.weight"),
3935+
3936+
# Shared Experts
3937+
f"{mt_layer_path}-mlp-shared_experts-wi_0-kernel": get_hf_key("mlp.shared_experts.gate_proj.weight"),
3938+
f"{mt_layer_path}-mlp-shared_experts-wi_1-kernel": get_hf_key("mlp.shared_experts.up_proj.weight"),
3939+
f"{mt_layer_path}-mlp-shared_experts-wo-kernel": get_hf_key("mlp.shared_experts.down_proj.weight"),
3940+
3941+
# Stacked Experts
3942+
f"{mt_layer_path}-mlp-MoeBlock_0-wi_0": get_hf_expert_keys("w1.weight"),
3943+
f"{mt_layer_path}-mlp-MoeBlock_0-wi_1": get_hf_expert_keys("w3.weight"),
3944+
f"{mt_layer_path}-mlp-MoeBlock_0-wo": get_hf_expert_keys("w2.weight"),
3945+
}
3946+
3947+
if (is_list and hf_layer_indices[0] >= 3) or (not is_list and hf_layer_indices >= 3):
3948+
layer_map[f"{mt_layer_path}-mlp-MoeBlock_0-gate-bias"] = get_hf_key("mlp.gate.e_score_correction_bias")
3949+
3950+
first_idx = hf_layer_indices[0] if is_list else hf_layer_indices
3951+
if first_idx >= 2:
3952+
if first_idx % 2 == 0 or first_idx == 2:
3953+
layer_map.update({
3954+
f"{mt_layer_path}-self_attention-csa_compressor-kv_proj-kernel": get_hf_key("self_attn.compressor.kv_proj.weight"),
3955+
f"{mt_layer_path}-self_attention-csa_compressor-gate_proj-kernel": get_hf_key("self_attn.compressor.gate_proj.weight"),
3956+
f"{mt_layer_path}-self_attention-csa_compressor-position_bias": get_hf_key("self_attn.compressor.position_bias"),
3957+
f"{mt_layer_path}-self_attention-csa_compressor-kv_norm-scale": get_hf_key("self_attn.compressor.kv_norm.weight"),
3958+
f"{mt_layer_path}-self_attention-csa_compressor-indexer-gate_proj-kernel": get_hf_key("self_attn.compressor.indexer.gate_proj.weight"),
3959+
f"{mt_layer_path}-self_attention-csa_compressor-indexer-kv_proj-kernel": get_hf_key("self_attn.compressor.indexer.kv_proj.weight"),
3960+
f"{mt_layer_path}-self_attention-csa_compressor-indexer-q_proj-kernel": get_hf_key("self_attn.compressor.indexer.q_b_proj.weight"),
3961+
f"{mt_layer_path}-self_attention-csa_compressor-indexer-weights_proj-kernel": get_hf_key("self_attn.compressor.indexer.scorer.weights_proj.weight"),
3962+
f"{mt_layer_path}-self_attention-csa_compressor-indexer-position_bias": get_hf_key("self_attn.compressor.indexer.position_bias"),
3963+
f"{mt_layer_path}-self_attention-csa_compressor-indexer-kv_norm-scale": get_hf_key("self_attn.compressor.indexer.kv_norm.weight"),
3964+
})
3965+
else:
3966+
layer_map.update({
3967+
f"{mt_layer_path}-self_attention-hca_compressor-kv_proj-kernel": get_hf_key("self_attn.compressor.kv_proj.weight"),
3968+
f"{mt_layer_path}-self_attention-hca_compressor-gate_proj-kernel": get_hf_key("self_attn.compressor.gate_proj.weight"),
3969+
f"{mt_layer_path}-self_attention-hca_compressor-position_bias": get_hf_key("self_attn.compressor.position_bias"),
3970+
f"{mt_layer_path}-self_attention-hca_compressor-kv_norm-scale": get_hf_key("self_attn.compressor.kv_norm.weight"),
3971+
})
3972+
3973+
mapping.update(layer_map)
3974+
3975+
if not scan_layers:
3976+
for i in range(n_layers):
3977+
add_layer_mapping(f"params-decoder-layers_{i}", i)
3978+
else:
3979+
for i in range(3):
3980+
add_layer_mapping(f"params-decoder-layers_{i}", i)
3981+
add_layer_mapping("params-decoder-scanned_blocks-layers_0", list(range(3, n_layers, 2)))
3982+
add_layer_mapping("params-decoder-scanned_blocks-layers_1", list(range(4, n_layers, 2)))
3983+
3984+
for i in range(3):
3985+
mapping[f"Tid2EidVar-decoder-layers_{i}-mlp-MoeBlock_0-tid2eid"] = f"model.layers.{i}.mlp.gate.tid2eid"
3986+
3987+
return mapping
3988+
3989+
def DEEPSEEKV4_MAXTEXT_TO_HF_PARAM_HOOK_FN(config, maxtext_config, scan_layers=False, saving_to_hf=False):
3990+
def transpose(input_tensor, target_shape=None):
3991+
return np.transpose(input_tensor)
3992+
3993+
def transpose_stack(input_tensor, target_shape=None):
3994+
# input_tensor is a list of tensors
3995+
stacked = np.stack(input_tensor, axis=0) # [E, out, in]
3996+
return np.transpose(stacked, (0, 2, 1)) # [E, in, out]
3997+
3998+
def ones_norm(input_tensor, target_shape=None):
3999+
return np.ones(target_shape, dtype=np.float32)
4000+
4001+
def identity(input_tensor, target_shape=None):
4002+
return input_tensor
4003+
4004+
# Reshaping functions for wq_b, wkv, o_a_proj
4005+
def reshape_transpose_wq_b(input_tensor, target_shape=None):
4006+
# HF: [n_heads * q_head_dim, kv_lora_rank]
4007+
# MaxText: [kv_lora_rank, n_heads, q_head_dim]
4008+
tensor = np.transpose(input_tensor) # [kv_lora_rank, n_heads * q_head_dim]
4009+
n_heads = config["num_attention_heads"]
4010+
return tensor.reshape(target_shape)
4011+
4012+
def reshape_transpose_wkv(input_tensor, target_shape=None):
4013+
# HF: [n_kv_heads * (q_head_dim + v_head_dim), kv_lora_rank]
4014+
# MaxText: [kv_lora_rank, n_kv_heads, q_head_dim + v_head_dim]
4015+
tensor = np.transpose(input_tensor)
4016+
return tensor.reshape(target_shape)
4017+
4018+
def reshape_transpose_o_a(input_tensor, target_shape=None):
4019+
# HF: [n_heads * v_head_dim, kv_lora_rank] (e.g. [8192, 4096])
4020+
# MaxText: [n_heads, v_head_dim, kv_lora_rank] (e.g. [8, 4096, 1024])
4021+
# We must reshape first and then permute (transpose) to get correct ordering.
4022+
num_heads = target_shape[0]
4023+
embed_dim = target_shape[1]
4024+
kv_lora_rank = target_shape[2]
4025+
tensor = input_tensor.reshape((num_heads, kv_lora_rank, embed_dim))
4026+
return np.transpose(tensor, (0, 2, 1))
4027+
4028+
# Functions for mHC split
4029+
def mhc_split_fn_pre(input_tensor, target_shape=None):
4030+
return np.transpose(input_tensor[0:4, :])
4031+
def mhc_split_fn_post(input_tensor, target_shape=None):
4032+
return np.transpose(input_tensor[4:8, :])
4033+
def mhc_split_fn_res(input_tensor, target_shape=None):
4034+
return np.transpose(input_tensor[8:24, :])
4035+
4036+
def mhc_split_base_pre(input_tensor, target_shape=None):
4037+
return input_tensor[0:4]
4038+
def mhc_split_base_post(input_tensor, target_shape=None):
4039+
return input_tensor[4:8]
4040+
def mhc_split_base_res(input_tensor, target_shape=None):
4041+
return input_tensor[8:24].reshape(target_shape)
4042+
4043+
def mhc_split_scale_pre(input_tensor, target_shape=None):
4044+
return np.array([input_tensor[0]]).reshape(target_shape)
4045+
def mhc_split_scale_post(input_tensor, target_shape=None):
4046+
return np.array([input_tensor[1]]).reshape(target_shape)
4047+
def mhc_split_scale_res(input_tensor, target_shape=None):
4048+
return np.array([input_tensor[2]]).reshape(target_shape)
4049+
4050+
mapping = {}
4051+
4052+
# Base mapping logic from original file
4053+
for key, hf_key in DEEPSEEKV4_MAXTEXT_TO_HF_PARAM_MAPPING(config, maxtext_config, scan_layers).items():
4054+
if hf_key is None:
4055+
mapping[key] = ones_norm
4056+
elif "token_embedder-embedding" in key:
4057+
mapping[key] = identity
4058+
elif "-wkv-kernel" in key:
4059+
mapping[key] = reshape_transpose_wkv
4060+
elif "-wq_b-kernel" in key:
4061+
mapping[key] = reshape_transpose_wq_b
4062+
elif "-o_a_proj-kernel" in key:
4063+
mapping[key] = reshape_transpose_o_a
4064+
elif "mhc" in key:
4065+
if "pre_alpha" in key and "scale" not in key: mapping[key] = mhc_split_fn_pre
4066+
elif "post_alpha" in key and "scale" not in key: mapping[key] = mhc_split_fn_post
4067+
elif "res_alpha" in key and "scale" not in key: mapping[key] = mhc_split_fn_res
4068+
elif "pre_beta" in key: mapping[key] = mhc_split_base_pre
4069+
elif "post_beta" in key: mapping[key] = mhc_split_base_post
4070+
elif "res_beta" in key: mapping[key] = mhc_split_base_res
4071+
elif "pre_alpha_scale" in key: mapping[key] = mhc_split_scale_pre
4072+
elif "post_alpha_scale" in key: mapping[key] = mhc_split_scale_post
4073+
elif "res_alpha_scale" in key: mapping[key] = mhc_split_scale_res
4074+
elif "position_bias" in key:
4075+
mapping[key] = identity
4076+
elif "hc_head-hc_fn" in key:
4077+
mapping[key] = transpose
4078+
elif "hc_head-hc_base" in key or "hc_head-hc_scale" in key:
4079+
mapping[key] = identity
4080+
elif type(hf_key) == list:
4081+
mapping[key] = transpose if not saving_to_hf else transpose_stack
4082+
elif "-kernel" in key or "-embedding" in key or "-sinks" in key:
4083+
mapping[key] = transpose
4084+
4085+
return mapping
4086+
38614087
PARAM_MAPPING = {
38624088
"gemma2-2b": GEMMA2_MAXTEXT_TO_HF_PARAM_MAPPING,
38634089
"gemma2-9b": GEMMA2_MAXTEXT_TO_HF_PARAM_MAPPING,
@@ -3896,6 +4122,7 @@ def reshape_vision_attn_out(input_tensor, target_shape):
38964122
"deepseek2-16b": DEEPSEEK_MAXTEXT_TO_HF_PARAM_MAPPING,
38974123
"deepseek3-671b": DEEPSEEK_MAXTEXT_TO_HF_PARAM_MAPPING,
38984124
"deepseek3.2-671b": DEEPSEEK_MAXTEXT_TO_HF_PARAM_MAPPING,
4125+
"deepseek4-284b": DEEPSEEKV4_MAXTEXT_TO_HF_PARAM_MAPPING,
38994126
"gpt-oss-20b": GPT_OSS_MAXTEXT_TO_HF_PARAM_MAPPING,
39004127
"gpt-oss-120b": GPT_OSS_MAXTEXT_TO_HF_PARAM_MAPPING,
39014128
"qwen3-omni-30b-a3b": QWEN3_OMNI_MOE_MAXTEXT_TO_HF_PARAM_MAPPING,
@@ -3948,6 +4175,8 @@ def reshape_vision_attn_out(input_tensor, target_shape):
39484175
"deepseek2-16b": DEEPSEEK_MAXTEXT_TO_HF_PARAM_HOOK_FN,
39494176
"deepseek3-671b": DEEPSEEK_MAXTEXT_TO_HF_PARAM_HOOK_FN,
39504177
"deepseek3.2-671b": DEEPSEEK_MAXTEXT_TO_HF_PARAM_HOOK_FN,
4178+
"deepseek4-tiny": DEEPSEEKV4_MAXTEXT_TO_HF_PARAM_HOOK_FN,
4179+
"deepseek4-284b": DEEPSEEKV4_MAXTEXT_TO_HF_PARAM_HOOK_FN,
39514180
"gpt-oss-20b": GPT_OSS_TO_HF_PARAM_HOOK_FN,
39524181
"gpt-oss-120b": GPT_OSS_TO_HF_PARAM_HOOK_FN,
39534182
"qwen3-omni-30b-a3b": QWEN3_OMNI_MOE_MAXTEXT_TO_HF_PARAM_HOOK_FN,

src/maxtext/configs/models/deepseek4-284b.yml

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -49,6 +49,10 @@ num_experts_per_tok: 6
4949
mlp_activations_limit: 10
5050
shared_experts: 1
5151
routed_score_func: "sqrtsoftplus"
52+
norm_topk_prob: true
53+
routed_bias: true
54+
routed_scaling_factor: 1.5
55+
5256

5357
# --- Attention configuration ---
5458
attention_type: 'compressed'
@@ -62,3 +66,4 @@ rope_type: "default"
6266
rope_max_timescale: 10000 # Main RoPE theta
6367
compressed_rope_max_timescale: 160000 # Compressed RoPE theta
6468
max_position_embeddings: 1048576
69+
original_max_position_embeddings: 65536
Lines changed: 71 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,71 @@
1+
# Copyright 2023–2026 Google LLC
2+
#
3+
# Licensed under the Apache License, Version 2.0 (the "License");
4+
# you may not use this file except in compliance with the License.
5+
# You may obtain a copy of the License at
6+
#
7+
# https://www.apache.org/licenses/LICENSE-2.0
8+
#
9+
# Unless required by applicable law or agreed to in writing, software
10+
# distributed under the License is distributed on an "AS IS" BASIS,
11+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+
# See the License for the specific language governing permissions and
13+
# limitations under the License.
14+
15+
# Model config for DeepSeek-V4-Flash 284B (https://huggingface.co/deepseek-ai/DeepSeek-V4-Flash)
16+
17+
base_emb_dim: 4096
18+
base_num_query_heads: 64
19+
base_num_kv_heads: 1
20+
base_num_decoder_layers: 7
21+
base_mlp_dim: 2048
22+
base_moe_mlp_dim: 2048
23+
vocab_size: 129280
24+
head_dim: 512
25+
26+
# --- Standard Defaults ---
27+
enable_dropout: false
28+
logits_via_embedding: false
29+
normalization_layer_epsilon: 1.0e-6
30+
31+
# --- V4 Specific Architectural Keys ---
32+
decoder_block: "deepseek4"
33+
mhc_expansion_rate: 4
34+
first_num_hash_layers: 3
35+
indexer_head_dim: 128
36+
indexer_n_heads: 64
37+
indexer_topk: 512
38+
39+
# Note: Layers (0, 1, 2) are prefix layers as `first_num_hash_layers=3`.
40+
# The 6th layer (MTP module with compress_ratio=0) has been explicitly dropped for now.
41+
# This leaves exactly 7 layers: 3 prefix [0,0,4] + 4 scanned.
42+
# `compress_ratio=0` uses sliding window attention. In this case, layer (0, 1).
43+
# This is a tiny version of deepseek4 with fewer layers and less experts for debugging.
44+
compress_ratios: [0, 0, 4, 128, 4, 128, 4]
45+
46+
# --- MoE configuration ---
47+
mlp_activations: ["silu", "linear"]
48+
num_experts: 8
49+
num_experts_per_tok: 3
50+
mlp_activations_limit: 10
51+
shared_experts: 1
52+
routed_score_func: "sqrtsoftplus"
53+
routed_bias: true
54+
routed_scaling_factor: 1.5
55+
56+
57+
# --- Attention configuration ---
58+
attention_type: 'compressed'
59+
attention: 'dot_product'
60+
q_lora_rank: 1024
61+
o_groups: 8
62+
o_lora_rank: 1024
63+
sliding_window_size: 128
64+
65+
# --- RoPE ---
66+
rope_type: "default"
67+
rope_max_timescale: 10000 # Main RoPE theta
68+
compressed_rope_max_timescale: 160000 # Compressed RoPE theta
69+
max_position_embeddings: 1048576
70+
original_max_position_embeddings: 65536
71+

src/maxtext/configs/types.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -228,6 +228,7 @@ class ProfilerType(str, Enum):
228228
"deepseek3-test",
229229
"deepseek3-tiny",
230230
"deepseek3.2-671b",
231+
"deepseek4-tiny",
231232
"deepseek4-284b",
232233
"deepseek-custom",
233234
"kimi-k2-1t",

src/maxtext/layers/attention_compressed.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -304,7 +304,7 @@ def __call__(
304304

305305
# Skip causal mask generation during decoding (seq_len == 1) or if no blocks were pooled
306306
if seq_len == 1 or compressed_len == 0:
307-
return compressed_kv, None
307+
return compressed_kv, jnp.zeros((batch_size, 1, seq_len, compressed_len), dtype=self.dtype)
308308

309309
# Construct a causal mask preventing early queries from attending to future compressed blocks
310310
entry_indices = jnp.arange(compressed_len)

src/maxtext/layers/decoders.py

Lines changed: 9 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1263,8 +1263,15 @@ def __call__(
12631263

12641264
# After the final transformer layer, `y` holds the raw, un-normalized hidden state.
12651265
if cfg.mhc_expansion_rate > 1:
1266-
# (batch, length, mhc_expansion_rate, emb_dim) --> (batch, length, emb_dim)
1267-
hidden_state = mhc_reduce(y)
1266+
if cfg.decoder_block == DecoderBlockType.DEEPSEEK4:
1267+
hidden_state = mhc.DeepSeek4HyperHeadToLinen(
1268+
config=cfg,
1269+
mesh=mesh,
1270+
name="hc_head",
1271+
)(y)
1272+
else:
1273+
# (batch, length, mhc_expansion_rate, emb_dim) --> (batch, length, emb_dim)
1274+
hidden_state = mhc_reduce(y)
12681275
else:
12691276
hidden_state = y
12701277

0 commit comments

Comments
 (0)