System Info
transformers version: 5.6.2
- Platform: Linux-6.8.0-1043-nvidia-x86_64-with-glibc2.35
- Python version: 3.12.13
- Huggingface_hub version: 1.11.0
- Safetensors version: 0.7.0
- Accelerate version: 1.13.0
- Accelerate config: not found
- DeepSpeed version: not installed
- PyTorch version (accelerator?): 2.10.0+cu129 (CUDA)
- Using distributed or parallel set-up in script?: no
- Using GPU in script?: yes
- GPU type: NVIDIA H100 80GB HBM3
Who can help?
@ArthurZucker @Cyrilvallez @3outeille
Information
Tasks
Reproduction
A 2-layer Gemma-4 from scratch (no checkpoint download), single GPU. Layer 1 is a KV-sharing layer that reads shared_kv_states[0] after layer 0 (a KV-source) was supposed to populate it.
"""`torchrun --nproc_per_node=1 repro.py` → KeyError(0)."""
import os
import torch
import torch.distributed as dist
from torch.distributed._composable.fsdp import MixedPrecisionPolicy, fully_shard
from torch.distributed.device_mesh import init_device_mesh
from transformers import Gemma4TextConfig, Gemma4TextModel
dist.init_process_group(backend="nccl")
torch.cuda.set_device(dist.get_rank())
cast = os.environ.get("CAST_FORWARD_INPUTS", "1") == "1"
print(f"cast_forward_inputs={cast}", flush=True)
# 2-layer Gemma-4: layer 0 writes shared_kv_states[0], layer 1 reads it.
# * num_kv_shared_layers=1 — default 0 means no sharing layer, bug can't fire.
# * layer_types must be set (default None breaks model init); both must be
# full_attention because Gemma-4 force-overrides the last layer to full,
# so the first must match for the sharing layer's same-type source lookup.
config = Gemma4TextConfig(
num_hidden_layers=2,
num_kv_shared_layers=1,
layer_types=["full_attention", "full_attention"],
)
model = Gemma4TextModel(config).cuda().train()
model.gradient_checkpointing_enable(
gradient_checkpointing_kwargs={"use_reentrant": True}
)
mesh = init_device_mesh("cuda", (dist.get_world_size(),))
mp = MixedPrecisionPolicy(param_dtype=torch.bfloat16, cast_forward_inputs=cast)
for layer in model.layers: # per-layer fully_shard triggers the per-call rebuild
fully_shard(layer, mesh=mesh, mp_policy=mp)
fully_shard(model, mesh=mesh, mp_policy=mp)
out = model(input_ids=torch.randint(0, config.vocab_size, (1, 8), device="cuda"))
print(f"PASS — last_hidden_state.shape={tuple(out.last_hidden_state.shape)}", flush=True)
dist.destroy_process_group()
Output:
cast_forward_inputs=True
File ".../transformers/models/gemma4/modeling_gemma4.py", line 1218, in forward
key_states, value_states = shared_kv_states[self.kv_shared_layer_index]
KeyError: 0
Expected behavior
Gemma4TextModel.forward creates shared_kv_states = {} once per forward (modeling_gemma4.py#L1668-L1669) and threads it through every decoder-layer call so later "sharing" layers can read earlier "source" layers' KV that were written by in-place mutation.
Under FSDP2 this breaks: each fully_shard-wrapped decoder layer's pre-forward traverses kwargs and rebuilds the inner dict. Layer 22's in-place write lands in an orphan; layer 25's read of shared_kv_states[22] raises KeyError: 22. The error message points at modeling_gemma4.py but the rebuild is entirely in the FSDP path.
Note that Gemma4TextModel already declares _skip_keys_device_placement = ["past_key_values", "shared_kv_states"] (modeling_gemma4.py#L1445), so this kwarg is special, but it's not used or passed to FSDP2.
Here's the stack trace I actually see:
File "transformers/modeling_layers.py", line 92, in __call__
return self._gradient_checkpointing_func(partial(super().__call__, **kwargs), *args)
File "torch/utils/checkpoint.py", line 268, in CheckpointFunction.forward
outputs = run_function(*args)
File "transformers/models/gemma4/modeling_gemma4.py", line 1219, in Gemma4TextAttention.forward
key_states, value_states = shared_kv_states[self.kv_shared_layer_index]
KeyError: 22
Per-layer dict-identity dump (one line per layer per rank, same forward) shows every layer sees a different dict_id:
GEMMA4_KV_DEBUG layer_idx=22 store_full_length_kv=True dict_id=15592302014016 dict_len=0
GEMMA4_KV_DEBUG layer_idx=24 is_kv_shared=True kv_shared_layer_index=22 dict_id=15592302338496 dict_len=0
GEMMA4_KV_DEBUG layer_idx=23 store_full_length_kv=True dict_id=11041820011008 dict_len=0
GEMMA4_KV_DEBUG layer_idx=24 is_kv_shared=True kv_shared_layer_index=22 dict_id=11041820015872 dict_len=0
Layer 22 is correctly configured (the upstream __init__ math says store_full_length_kv=True), but its mutation is lost because the dict it writes to is not the same dict layer 24/25 reads from.
Root cause
_FSDPState._pre_forward (torch/distributed/fsdp/_fully_shard/_fsdp_state.py#L243-L251) under cast_forward_inputs=True:
args, kwargs = (
_apply_to_tensors(cast_fn, args),
_apply_to_tensors(cast_fn, kwargs),
)
_apply_to_tensors (torch/distributed/utils.py#L245-L246) rebuilds every dict it recurses into via {k: apply(v) for k, v in x.items()} — even when no tensor inside actually needs casting. So every per-decoder-layer wrapper hands the layer a freshly-allocated shared_kv_states... layer 22 mutates one orphan, layer 24 reads another.
System Info
transformersversion: 5.6.2Who can help?
@ArthurZucker @Cyrilvallez @3outeille
Information
Tasks
examplesfolder (such as GLUE/SQuAD, ...)Reproduction
A 2-layer Gemma-4 from scratch (no checkpoint download), single GPU. Layer 1 is a KV-sharing layer that reads
shared_kv_states[0]after layer 0 (a KV-source) was supposed to populate it.Output:
Expected behavior
Gemma4TextModel.forwardcreatesshared_kv_states = {}once per forward (modeling_gemma4.py#L1668-L1669) and threads it through every decoder-layer call so later "sharing" layers can read earlier "source" layers' KV that were written by in-place mutation.Under FSDP2 this breaks: each
fully_shard-wrapped decoder layer's pre-forward traverses kwargs and rebuilds the inner dict. Layer 22's in-place write lands in an orphan; layer 25's read ofshared_kv_states[22]raisesKeyError: 22. The error message points atmodeling_gemma4.pybut the rebuild is entirely in the FSDP path.Note that
Gemma4TextModelalready declares_skip_keys_device_placement = ["past_key_values", "shared_kv_states"](modeling_gemma4.py#L1445), so this kwarg is special, but it's not used or passed to FSDP2.Here's the stack trace I actually see:
Per-layer dict-identity dump (one line per layer per rank, same forward) shows every layer sees a different
dict_id:Layer 22 is correctly configured (the upstream
__init__math saysstore_full_length_kv=True), but its mutation is lost because the dict it writes to is not the same dict layer 24/25 reads from.Root cause
_FSDPState._pre_forward(torch/distributed/fsdp/_fully_shard/_fsdp_state.py#L243-L251) undercast_forward_inputs=True:_apply_to_tensors(torch/distributed/utils.py#L245-L246) rebuilds every dict it recurses into via{k: apply(v) for k, v in x.items()}— even when no tensor inside actually needs casting. So every per-decoder-layer wrapper hands the layer a freshly-allocatedshared_kv_states... layer 22 mutates one orphan, layer 24 reads another.