def forward(self, x, kv_cache=None, start_pos=0): T_curr = x.size(1) position_ids = torch.arange(start_pos, start_pos + T_curr, device=x.device) cos, sin = self.rotary_embd(position_ids) for i, block in enumerate(self.blocks): # Pass per-layer KV cache x, kv_cache[i] = block(x, cos, sin, attention_mask, kv_cache[i]) return x, kv_cache