Download src/jolt/execution.py from mlengineer-ai/Jolt-2B: direct link, hf CLI and curl.
- Browser
- Download file 5.31 kB
-
https://huggingface.co/mlengineer-ai/Jolt-2B/resolve/main/src/jolt/execution.py
- Command line
-
hf download hf://mlengineer-ai/Jolt-2B/src/jolt/execution.py
-
curl -L -o execution.py https://huggingface.co/mlengineer-ai/Jolt-2B/resolve/main/src/jolt/execution.py
5.31 kB
| """Request-scoped, differentiable prefix sharing for the Qwen hybrid cache.""" | |
| import torch | |
| from transformers.cache_utils import DynamicCache | |
| def plan_prefix(inputs, positions): | |
| """Plan on CPU, before asynchronous transfer; candidates stay in each branch.""" | |
| ids = inputs["input_ids"] | |
| cut = 0 | |
| images = inputs.get("image_keys", ()) | |
| if len(ids) > 1 and len(set(images)) <= 1: | |
| common = (ids == ids[:1]).all(0) & inputs["attention_mask"].bool().all(0) | |
| different = (~common).nonzero() | |
| length = int(different[0]) if len(different) else common.numel() | |
| cut = min(length, int(positions[:, 0].min()), ids.shape[1] - 2) // 64 * 64 | |
| inputs["prefix_length"] = cut | |
| class PrefixCache(DynamicCache): | |
| """Replace state tensors instead of overwriting values needed by backward.""" | |
| def update_conv_state( | |
| self, conv_states, layer_idx, state_idx=0, conv_kernel_size=None, **kwargs | |
| ): | |
| if conv_kernel_size is None: | |
| raise ValueError("PrefixCache requires conv_kernel_size") | |
| layer = self.layers[layer_idx] | |
| layer.device, layer.dtype = conv_states.device, conv_states.dtype | |
| if layer.has_previous_state[state_idx]: | |
| conv_states = torch.cat([layer.conv_states[state_idx], conv_states], dim=-1) | |
| layer.conv_kernel_size[state_idx] = conv_kernel_size | |
| layer.conv_states[state_idx] = conv_states[..., -conv_kernel_size:] | |
| layer.has_previous_state[state_idx] = True | |
| layer.is_conv_states_initialized[state_idx] = True | |
| return conv_states | |
| def update_recurrent_state(self, recurrent_states, layer_idx, state_idx=0, **kwargs): | |
| layer = self.layers[layer_idx] | |
| layer.recurrent_states[state_idx] = recurrent_states | |
| layer.is_recurrent_states_initialized[state_idx] = True | |
| return recurrent_states | |
| def language_forward( | |
| language, | |
| embeddings, | |
| attention_mask, | |
| position_ids, | |
| cut, | |
| capture_layers=(), | |
| capture_raw_layers=None, | |
| ): | |
| requested = list(capture_layers) if capture_layers else False | |
| capture_raw_layers = capture_raw_layers or {} | |
| def forward( | |
| inputs_embeds, | |
| attention_mask, | |
| position_ids, | |
| *, | |
| past_key_values=None, | |
| use_cache=False, | |
| capture=True, | |
| ): | |
| return language( | |
| inputs_embeds=inputs_embeds, | |
| attention_mask=attention_mask, | |
| position_ids=position_ids, | |
| past_key_values=past_key_values, | |
| use_cache=use_cache, | |
| output_hidden_states=requested if capture else False, | |
| return_dict=True, | |
| ) | |
| def result(outputs, raw=None): | |
| if not capture_layers and not capture_raw_layers: | |
| return outputs.last_hidden_state | |
| raw = raw or {} | |
| return ( | |
| outputs.last_hidden_state, | |
| *(outputs.hidden_states[index] for index in capture_layers), | |
| *(raw[index] for index in capture_raw_layers), | |
| ) | |
| def captured_forward(inputs_embeds, attention_mask, position_ids, **kwargs): | |
| raw = {} | |
| handles = [] | |
| if kwargs.pop("capture", True): | |
| for index, layer in capture_raw_layers.items(): | |
| def save_raw(_module, _args, output, layer_index=index): | |
| raw[layer_index] = output[0] if isinstance(output, tuple) else output | |
| handles.append(layer.register_forward_hook(save_raw)) | |
| try: | |
| outputs = forward(inputs_embeds, attention_mask, position_ids, **kwargs) | |
| finally: | |
| for handle in handles: | |
| handle.remove() | |
| return result(outputs, raw) | |
| if not cut: | |
| return captured_forward(embeddings, attention_mask, position_ids) | |
| def shared_forward(): | |
| # A mutable cache cannot survive a decoder-layer checkpoint replay. | |
| # Replay the complete request instead, constructing a fresh cache each time. | |
| layers = [m for m in language.modules() if getattr(m, "gradient_checkpointing", False)] | |
| for layer in layers: | |
| layer.gradient_checkpointing = False | |
| try: | |
| cache = PrefixCache(config=language.config) | |
| forward( | |
| embeddings[:1, :cut], | |
| attention_mask[:1, :cut], | |
| None if position_ids is None else position_ids[:, :1, :cut], | |
| past_key_values=cache, | |
| use_cache=True, | |
| capture=False, | |
| ) | |
| # Transformers expects LongTensor; factories return Tensor with int64 dtype. | |
| indices = torch.zeros(len(embeddings), dtype=torch.long, device=embeddings.device) | |
| cache.reorder_cache(indices) # ty: ignore[invalid-argument-type] | |
| return captured_forward( | |
| embeddings[:, cut:], | |
| attention_mask, | |
| None if position_ids is None else position_ids[:, :, cut:], | |
| past_key_values=cache, | |
| use_cache=True, | |
| ) | |
| finally: | |
| for layer in layers: | |
| layer.gradient_checkpointing = True | |
| if torch.is_grad_enabled(): | |
| from torch.utils.checkpoint import checkpoint | |
| return checkpoint(shared_forward, use_reentrant=False) | |
| return shared_forward() | |