File size: 5,307 Bytes
f55039b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 | """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()
|