"""CUDA graph replay of Qwen3.5's unchanged recurrent decode operations.""" from __future__ import annotations import torch from transformers.cache_utils import DynamicCache from transformers.models.qwen3_5.modeling_qwen3_5 import Qwen3_5DecoderLayer class EdgeInstantDecoderLayer(Qwen3_5DecoderLayer): def __init__(self, config, layer_idx): super().__init__(config, layer_idx) self.config = config self.layer_idx = layer_idx self._decode_graph = None def train(self, mode=True): self._decode_graph = None return super().train(mode) def _apply(self, fn, recurse=True): self._decode_graph = None return super()._apply(fn, recurse=recurse) def forward(self, hidden_states, position_embeddings, attention_mask=None, position_ids=None, past_key_values=None, **kwargs): cache_params = past_key_values if (self.training or torch.is_grad_enabled() or hidden_states.device.type != "cuda" or hidden_states.shape[1] != 1 or cache_params is None or attention_mask is not None or torch.compiler.is_compiling() or not cache_params.has_previous_state(self.layer_idx)): return super().forward(hidden_states, position_embeddings, attention_mask, position_ids, past_key_values, **kwargs) layer = cache_params.layers[self.layer_idx] shape = (hidden_states.shape, hidden_states.dtype, hidden_states.device, layer.conv_states.shape, layer.conv_states.dtype, layer.recurrent_states.shape, layer.recurrent_states.dtype) state = self._decode_graph if state is None or state[0] != shape: inputs = torch.empty_like(hidden_states) graph_cache = DynamicCache(config=self.config) graph_cache.update_conv_state(layer.conv_states, self.layer_idx) graph_cache.update_recurrent_state(layer.recurrent_states, self.layer_idx) inputs.copy_(hidden_states) current = torch.cuda.current_stream(hidden_states.device) stream = torch.cuda.Stream(device=hidden_states.device) stream.wait_stream(current) with torch.cuda.stream(stream): for _ in range(3): super().forward(inputs, None, past_key_values=graph_cache) current.wait_stream(stream) graph = torch.cuda.CUDAGraph() with torch.cuda.graph(graph, stream=stream): outputs = super().forward(inputs, None, past_key_values=graph_cache) state = (shape, inputs, graph_cache, graph, outputs) self._decode_graph = state _, inputs, graph_cache, graph, outputs = state # Keep each caller's cache independent, including caches returned by generate. inputs.copy_(hidden_states) graph_cache.update_conv_state(layer.conv_states, self.layer_idx) graph_cache.update_recurrent_state(layer.recurrent_states, self.layer_idx) graph.replay() graph_layer = graph_cache.layers[self.layer_idx] cache_params.update_conv_state(graph_layer.conv_states, self.layer_idx) cache_params.update_recurrent_state(graph_layer.recurrent_states, self.layer_idx) return outputs.clone()