diff --git a/gemma/tensor_info.cc b/gemma/tensor_info.cc index 15c3517b..9841b5a3 100644 --- a/gemma/tensor_info.cc +++ b/gemma/tensor_info.cc @@ -741,7 +741,7 @@ void TensorInfoRegistry::AddLayerTensors(const ModelConfig& config, }); Add(suffix, { .base_name = "qkv1_w", - .source_names = {"attn/q_einsum/w"}, + .source_names = {"attn/q_einsum/w", "attn/q_einsum"}, .axes = {0, 2, 1}, .shape = {layer_config.heads * layer_config.qkv_dim, config.model_dim}, @@ -750,11 +750,18 @@ void TensorInfoRegistry::AddLayerTensors(const ModelConfig& config, Add(suffix, { .base_name = "qkv2_w", - .source_names = {layer_config.kv_heads == 1 - ? "attn/k_einsum/w" - : "attn/kv_einsum/w"}, - .axes = layer_config.kv_heads == 1 ? std::vector{0, 2, 1} - : std::vector{1, 0, 3, 2}, + .source_names = (layer_config.kv_heads == 1 && + config.model_family_version < 3) + ? std::vector{"attn/k_einsum/w", + "attn/k_einsum"} + : std::vector{"attn/kv_einsum/w", + "attn/kv_einsum", + "attn/k_einsum/w", + "attn/k_einsum"}, + .axes = (layer_config.kv_heads == 1 && + config.model_family_version < 3) + ? std::vector{0, 2, 1} + : std::vector{1, 0, 3, 2}, .shape = {2 * layer_config.kv_heads * layer_config.qkv_dim, config.model_dim}, .concat_names = {""}, @@ -1056,7 +1063,7 @@ void TensorInfoRegistry::AddLayerTensors(const ModelConfig& config, Add(suffix, { .base_name = "att_ein", - .source_names = {"attn/attn_vec_einsum/w", + .source_names = {"attn/attn_vec_einsum/w", "attn/attn_vec_einsum", "attention_block/proj_final/kernel"}, .preshape = {layer_config.heads, layer_config.qkv_dim, config.model_dim}, diff --git a/gemma/weights.h b/gemma/weights.h index 39c1c4a8..2a644db1 100644 --- a/gemma/weights.h +++ b/gemma/weights.h @@ -463,11 +463,11 @@ struct LayerWeightsPtrs { func(TENSOR_ARGS(hash_tid2eid, kMustRead)); } } else { - func(TENSOR_ARGS(router_scale, kMustRead)); - func(TENSOR_ARGS(p_expert_sc, kMustRead)); - func(TENSOR_ARGS(post_ffw1_ns, kMustRead)); - func(TENSOR_ARGS(post_ffw2_ns, kMustRead)); - func(TENSOR_ARGS(pre_ffw2_ns, kMustRead)); + func(TENSOR_ARGS(router_scale, kMaybeRead)); + func(TENSOR_ARGS(p_expert_sc, kMaybeRead)); + func(TENSOR_ARGS(post_ffw1_ns, kMaybeRead)); + func(TENSOR_ARGS(post_ffw2_ns, kMaybeRead)); + func(TENSOR_ARGS(pre_ffw2_ns, kMaybeRead)); } for (uint32_t i = 0; i < layer_config.NumExperts(); ++i) { func(TENSOR_ARGS(moe_gating_einsum_w1[i], kMustRead));