Jolt-2B / src /jolt /execution.py
mehran's picture
Upload folder using huggingface_hub
f55039b verified
Raw History Blame Contribute Delete
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()