HPD-Parsing / modeling_internvl_chat.py
WEISHU's picture
Upload folder using huggingface_hub
7325252 verified
Raw
History Blame Contribute Delete
35.4 kB
# --------------------------------------------------------
# InternVL
# Copyright (c) 2024 OpenGVLab
# Licensed under The MIT License [see LICENSE for details]
# --------------------------------------------------------
import os
import warnings
from typing import List, Optional, Tuple, Union
import torch.utils.checkpoint
import torch.nn.functional as F
import transformers
from torch import nn
from torch.nn import CrossEntropyLoss
from transformers import GenerationConfig
from transformers.cache_utils import DynamicCache
from transformers.modeling_outputs import CausalLMOutputWithPast
from transformers.modeling_utils import PreTrainedModel
from transformers.utils import logging
from transformers import LlamaForCausalLM, Qwen2ForCausalLM, Qwen3ForCausalLM, Qwen3MoeForCausalLM
from .configuration_internvl_chat import InternVLChatConfig
from .conversation import get_conv_template
from .modeling_intern_vit import InternVisionModel, has_flash_attn
logger = logging.get_logger(__name__)
def version_cmp(v1, v2, op='eq'):
import operator
from packaging import version
op_func = getattr(operator, op)
return op_func(version.parse(v1), version.parse(v2))
class HPDRMSNorm(nn.Module):
def __init__(self, hidden_size, eps=1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(hidden_size))
self.variance_epsilon = eps
def forward(self, hidden_states):
input_dtype = hidden_states.dtype
hidden_states = hidden_states.to(torch.float32)
variance = hidden_states.pow(2).mean(-1, keepdim=True)
hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
return self.weight * hidden_states.to(input_dtype)
class _HPDGatedMLP(nn.Module):
def __init__(self, hidden_size, intermediate_size):
super().__init__()
self.gate_proj = nn.Linear(hidden_size, intermediate_size, bias=False)
self.up_proj = nn.Linear(hidden_size, intermediate_size, bias=False)
self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=False)
def forward(self, x):
return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))
class _HPDMTPLayer(nn.Module):
def __init__(self, hidden_size, intermediate_size):
super().__init__()
self.mlp = _HPDGatedMLP(hidden_size, intermediate_size)
class HPDMultiTokenPredictor(nn.Module):
"""Single-block Progressive Multi-Token Prediction head (P-MTP).
Mirrors the vLLM ``Qwen3_5MultiTokenPredictor`` module (loaded there as a
``medusa`` speculative model) so its weights, stored under
``language_model.mtp.*`` in the checkpoint, load without remapping. One MTP
block is applied autoregressively to draft several future tokens; each step
is conditioned on the previous step's predicted token. The token embedding
and the output head are shared with the backbone (the checkpoint contains no
``mtp.embed_tokens``/``mtp.lm_head``), so ``step`` takes an already-embedded
``prev_embeds`` tensor.
"""
def __init__(self, hidden_size, intermediate_size, rms_norm_eps=1e-6):
super().__init__()
self.fc = nn.Linear(hidden_size * 2, hidden_size, bias=False)
self.pre_fc_norm_hidden = HPDRMSNorm(hidden_size, rms_norm_eps)
self.pre_fc_norm_embedding = HPDRMSNorm(hidden_size, rms_norm_eps)
self.layers = nn.ModuleList([_HPDMTPLayer(hidden_size, intermediate_size)])
self.norm = HPDRMSNorm(hidden_size, rms_norm_eps)
def step(self, hidden_states, prev_embeds):
# hidden_states: [B, H] backbone hidden that produced prev_token;
# prev_embeds: [B, H] backbone embedding of that token.
embeds = self.pre_fc_norm_embedding(prev_embeds)
hidden = self.pre_fc_norm_hidden(hidden_states)
x = self.fc(torch.cat([hidden, embeds], dim=-1))
mlp_out = self.layers[0].mlp(x)
return self.norm(mlp_out + x)
class InternVLChatModel(PreTrainedModel):
config_class = InternVLChatConfig
main_input_name = 'pixel_values'
base_model_prefix = 'language_model'
_supports_flash_attn_2 = True
supports_gradient_checkpointing = True
_no_split_modules = [
"InternVisionModel",
"Qwen3DecoderLayer",
]
# support transformers 4.51.+
_tp_plan = ''
def __init__(self, config: InternVLChatConfig, vision_model=None, language_model=None, use_flash_attn=True):
super().__init__(config)
assert version_cmp(transformers.__version__, '4.37.0', 'ge')
image_size = config.force_image_size or config.vision_config.image_size
patch_size = config.vision_config.patch_size
self.patch_size = patch_size
self.select_layer = config.select_layer
self.template = config.template
self.num_image_token = int((image_size // patch_size) ** 2 * (config.downsample_ratio ** 2))
self.downsample_ratio = config.downsample_ratio
self.ps_version = config.ps_version
use_flash_attn = use_flash_attn if has_flash_attn else False
config.vision_config.use_flash_attn = True if use_flash_attn else False
config.llm_config._attn_implementation = 'flash_attention_2' if use_flash_attn else 'eager'
logger.info(f'num_image_token: {self.num_image_token}')
logger.info(f'ps_version: {self.ps_version}')
if vision_model is not None:
self.vision_model = vision_model
else:
self.vision_model = InternVisionModel(config.vision_config)
if language_model is not None:
self.language_model = language_model
else:
architecture: str = config.llm_config.architectures[0]
if architecture == 'LlamaForCausalLM':
self.language_model = LlamaForCausalLM(config.llm_config)
elif architecture == 'Qwen2ForCausalLM':
self.language_model = Qwen2ForCausalLM(config.llm_config)
elif architecture == 'Qwen3MoeForCausalLM':
self.language_model = Qwen3MoeForCausalLM(config.llm_config)
elif architecture == 'Qwen3ForCausalLM':
self.language_model = Qwen3ForCausalLM(config.llm_config)
else:
raise NotImplementedError(f'{architecture} is not implemented.')
vit_hidden_size = config.vision_config.hidden_size
llm_hidden_size = config.llm_config.hidden_size
self.mlp1 = nn.Sequential(
nn.LayerNorm(vit_hidden_size * int(1 / self.downsample_ratio) ** 2),
nn.Linear(vit_hidden_size * int(1 / self.downsample_ratio) ** 2, llm_hidden_size),
nn.GELU(),
nn.Linear(llm_hidden_size, llm_hidden_size)
)
self.img_context_token_id = None
self.conv_template = get_conv_template(self.template)
self.system_message = self.conv_template.system_message
# P-MTP head + hierarchical fork decoding metadata (see generate_hpd).
self.fork_token_id = getattr(config, 'fork_token_id', None)
self.child_token_id = getattr(config, 'child_token_id', None)
self._mtp_weights_loaded = False
llm_cfg = config.llm_config
self.language_model.mtp = HPDMultiTokenPredictor(
hidden_size=llm_cfg.hidden_size,
intermediate_size=llm_cfg.intermediate_size,
rms_norm_eps=getattr(llm_cfg, 'rms_norm_eps', 1e-6),
)
def forward(
self,
pixel_values: torch.FloatTensor,
input_ids: torch.LongTensor = None,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.LongTensor] = None,
image_flags: Optional[torch.LongTensor] = None,
past_key_values: Optional[List[torch.FloatTensor]] = None,
labels: Optional[torch.LongTensor] = None,
use_cache: Optional[bool] = None,
output_attentions: Optional[bool] = None,
output_hidden_states: Optional[bool] = None,
return_dict: Optional[bool] = None,
) -> Union[Tuple, CausalLMOutputWithPast]:
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
image_flags = image_flags.squeeze(-1)
input_embeds = self.language_model.get_input_embeddings()(input_ids).clone()
vit_embeds = self.extract_feature(pixel_values)
vit_embeds = vit_embeds[image_flags == 1]
vit_batch_size = pixel_values.shape[0]
B, N, C = input_embeds.shape
input_embeds = input_embeds.reshape(B * N, C)
# if torch.distributed.is_initialized() and torch.distributed.get_rank() == 0:
# print(f'dynamic ViT batch size: {vit_batch_size}, images per sample: {vit_batch_size / B}, dynamic token length: {N}')
input_ids = input_ids.reshape(B * N)
selected = (input_ids == self.img_context_token_id)
try:
input_embeds[selected] = input_embeds[selected] * 0.0 + vit_embeds.reshape(-1, C)
except Exception as e:
vit_embeds = vit_embeds.reshape(-1, C)
print(f'warning: {e}, input_embeds[selected].shape={input_embeds[selected].shape}, '
f'vit_embeds.shape={vit_embeds.shape}')
n_token = min(selected.sum(), vit_embeds.size(0))
input_embeds[selected][:n_token] = input_embeds[selected][:n_token] * 0.0 + vit_embeds[:n_token]
input_embeds = input_embeds.reshape(B, N, C)
outputs = self.language_model(
inputs_embeds=input_embeds,
attention_mask=attention_mask,
position_ids=position_ids,
past_key_values=past_key_values,
use_cache=use_cache,
output_attentions=output_attentions,
output_hidden_states=output_hidden_states,
return_dict=return_dict,
)
logits = outputs.logits
loss = None
if labels is not None:
# Shift so that tokens < n predict n
shift_logits = logits[..., :-1, :].contiguous()
shift_labels = labels[..., 1:].contiguous()
# Flatten the tokens
loss_fct = CrossEntropyLoss()
shift_logits = shift_logits.view(-1, self.language_model.config.vocab_size)
shift_labels = shift_labels.view(-1)
# Enable model parallelism
shift_labels = shift_labels.to(shift_logits.device)
loss = loss_fct(shift_logits, shift_labels)
if not return_dict:
output = (logits,) + outputs[1:]
return (loss,) + output if loss is not None else output
return CausalLMOutputWithPast(
loss=loss,
logits=logits,
past_key_values=outputs.past_key_values,
hidden_states=outputs.hidden_states,
attentions=outputs.attentions,
)
def pixel_shuffle(self, x, scale_factor=0.5):
n, w, h, c = x.size()
# N, W, H, C --> N, W, H * scale, C // scale
x = x.view(n, w, int(h * scale_factor), int(c / scale_factor))
# N, W, H * scale, C // scale --> N, H * scale, W, C // scale
x = x.permute(0, 2, 1, 3).contiguous()
# N, H * scale, W, C // scale --> N, H * scale, W * scale, C // (scale ** 2)
x = x.view(n, int(h * scale_factor), int(w * scale_factor),
int(c / (scale_factor * scale_factor)))
if self.ps_version == 'v1':
warnings.warn("In ps_version 'v1', the height and width have not been swapped back, "
'which results in a transposed image.')
else:
x = x.permute(0, 2, 1, 3).contiguous()
return x
def extract_feature(self, pixel_values):
if self.select_layer == -1:
vit_embeds = self.vision_model(
pixel_values=pixel_values,
output_hidden_states=False,
return_dict=True).last_hidden_state
else:
vit_embeds = self.vision_model(
pixel_values=pixel_values,
output_hidden_states=True,
return_dict=True).hidden_states[self.select_layer]
vit_embeds = vit_embeds[:, 1:, :]
h = w = int(vit_embeds.shape[1] ** 0.5)
vit_embeds = vit_embeds.reshape(vit_embeds.shape[0], h, w, -1)
vit_embeds = self.pixel_shuffle(vit_embeds, scale_factor=self.downsample_ratio)
vit_embeds = vit_embeds.reshape(vit_embeds.shape[0], -1, vit_embeds.shape[-1])
vit_embeds = self.mlp1(vit_embeds)
return vit_embeds
def batch_chat(self, tokenizer, pixel_values, questions, generation_config, num_patches_list=None,
history=None, return_history=False, IMG_START_TOKEN='<img>', IMG_END_TOKEN='</img>',
IMG_CONTEXT_TOKEN='<IMG_CONTEXT>', verbose=False, image_counts=None):
if history is not None or return_history:
print('Now multi-turn chat is not supported in batch_chat.')
raise NotImplementedError
if image_counts is not None:
num_patches_list = image_counts
print('Warning: `image_counts` is deprecated. Please use `num_patches_list` instead.')
img_context_token_id = tokenizer.convert_tokens_to_ids(IMG_CONTEXT_TOKEN)
self.img_context_token_id = img_context_token_id
if verbose and pixel_values is not None:
image_bs = pixel_values.shape[0]
print(f'dynamic ViT batch size: {image_bs}')
queries = []
for idx, num_patches in enumerate(num_patches_list):
question = questions[idx]
if pixel_values is not None and '<image>' not in question:
question = '<image>\n' + question
template = get_conv_template(self.template)
template.system_message = self.system_message
template.append_message(template.roles[0], question)
template.append_message(template.roles[1], None)
query = template.get_prompt()
image_tokens = IMG_START_TOKEN + IMG_CONTEXT_TOKEN * self.num_image_token * num_patches + IMG_END_TOKEN
query = query.replace('<image>', image_tokens, 1)
queries.append(query)
tokenizer.padding_side = 'left'
model_inputs = tokenizer(queries, return_tensors='pt', padding=True)
input_ids = model_inputs['input_ids'].to(self.device)
attention_mask = model_inputs['attention_mask'].to(self.device)
eos_token_id = tokenizer.convert_tokens_to_ids(template.sep.strip())
generation_config['eos_token_id'] = eos_token_id
generation_output = self.generate(
pixel_values=pixel_values,
input_ids=input_ids,
attention_mask=attention_mask,
**generation_config
)
responses = tokenizer.batch_decode(generation_output, skip_special_tokens=True)
responses = [response.split(template.sep.strip())[0].strip() for response in responses]
return responses
def chat(self, tokenizer, pixel_values, question, generation_config, history=None, return_history=False,
num_patches_list=None, IMG_START_TOKEN='<img>', IMG_END_TOKEN='</img>', IMG_CONTEXT_TOKEN='<IMG_CONTEXT>',
verbose=False):
if history is None and pixel_values is not None and '<image>' not in question:
question = '<image>\n' + question
if num_patches_list is None:
num_patches_list = [pixel_values.shape[0]] if pixel_values is not None else []
assert pixel_values is None or len(pixel_values) == sum(num_patches_list)
img_context_token_id = tokenizer.convert_tokens_to_ids(IMG_CONTEXT_TOKEN)
self.img_context_token_id = img_context_token_id
template = get_conv_template(self.template)
template.system_message = self.system_message
eos_token_id = tokenizer.convert_tokens_to_ids(template.sep.strip())
history = [] if history is None else history
for (old_question, old_answer) in history:
template.append_message(template.roles[0], old_question)
template.append_message(template.roles[1], old_answer)
template.append_message(template.roles[0], question)
template.append_message(template.roles[1], None)
query = template.get_prompt()
if verbose and pixel_values is not None:
image_bs = pixel_values.shape[0]
print(f'dynamic ViT batch size: {image_bs}')
for num_patches in num_patches_list:
image_tokens = IMG_START_TOKEN + IMG_CONTEXT_TOKEN * self.num_image_token * num_patches + IMG_END_TOKEN
query = query.replace('<image>', image_tokens, 1)
model_inputs = tokenizer(query, return_tensors='pt')
input_ids = model_inputs['input_ids'].to(self.device)
attention_mask = model_inputs['attention_mask'].to(self.device)
generation_config['eos_token_id'] = eos_token_id
generation_output = self.generate(
pixel_values=pixel_values,
input_ids=input_ids,
attention_mask=attention_mask,
**generation_config
)
response = tokenizer.batch_decode(generation_output, skip_special_tokens=True)[0]
response = response.split(template.sep.strip())[0].strip()
history.append((question, response))
if return_history:
return response, history
else:
query_to_print = query.replace(IMG_CONTEXT_TOKEN, '')
query_to_print = query_to_print.replace(f'{IMG_START_TOKEN}{IMG_END_TOKEN}', '<image>')
if verbose:
print(query_to_print, response)
return response
@torch.no_grad()
def generate(
self,
pixel_values: Optional[torch.FloatTensor] = None,
input_ids: Optional[torch.FloatTensor] = None,
attention_mask: Optional[torch.LongTensor] = None,
visual_features: Optional[torch.FloatTensor] = None,
generation_config: Optional[GenerationConfig] = None,
output_hidden_states: Optional[bool] = None,
**generate_kwargs,
) -> torch.LongTensor:
assert self.img_context_token_id is not None
if pixel_values is not None:
if visual_features is not None:
vit_embeds = visual_features
else:
vit_embeds = self.extract_feature(pixel_values)
input_embeds = self.language_model.get_input_embeddings()(input_ids)
B, N, C = input_embeds.shape
input_embeds = input_embeds.reshape(B * N, C)
input_ids = input_ids.reshape(B * N)
selected = (input_ids == self.img_context_token_id)
assert selected.sum() != 0
input_embeds[selected] = vit_embeds.reshape(-1, C).to(input_embeds.device)
input_embeds = input_embeds.reshape(B, N, C)
else:
input_embeds = self.language_model.get_input_embeddings()(input_ids)
outputs = self.language_model.generate(
inputs_embeds=input_embeds,
attention_mask=attention_mask,
generation_config=generation_config,
output_hidden_states=output_hidden_states,
use_cache=True,
**generate_kwargs,
)
return outputs
def load_mtp_weights(self, path=None, prefix='language_model.mtp.'):
"""Load P-MTP head weights.
If ``path`` is None the weights are assumed to already be present (e.g.
loaded from the main checkpoint via ``from_pretrained`` because the keys
``language_model.mtp.*`` match this module tree) and this call only marks
MTP as usable. Otherwise ``path`` may be a ``.safetensors`` file or a
directory containing one; keys are stripped of ``prefix`` (or a leading
``mtp.``) before loading.
"""
if path is None:
self._mtp_weights_loaded = True
return
from safetensors.torch import load_file
if os.path.isdir(path):
files = [os.path.join(path, f) for f in sorted(os.listdir(path))
if f.endswith('.safetensors')]
else:
files = [path]
remapped = {}
for fp in files:
for k, v in load_file(fp).items():
name = k
if name.startswith(prefix):
name = name[len(prefix):]
elif name.startswith('mtp.'):
name = name[len('mtp.'):]
else:
continue
remapped[name] = v
if not remapped:
raise RuntimeError(f'No MTP tensors (prefix {prefix!r}) found under {path}')
missing, unexpected = self.language_model.mtp.load_state_dict(remapped, strict=False)
self.language_model.mtp.to(device=self.device, dtype=self.dtype)
self._mtp_weights_loaded = True
return missing, unexpected
def _embed_ids(self, input_ids, vit_embeds=None):
embeds = self.language_model.get_input_embeddings()(input_ids)
if vit_embeds is not None:
B, N, C = embeds.shape
flat = embeds.reshape(B * N, C)
selected = (input_ids.reshape(B * N) == self.img_context_token_id)
flat[selected] = vit_embeds.reshape(-1, C).to(flat.dtype)
embeds = flat.reshape(B, N, C)
return embeds
def _lm_forward(self, inputs_embeds, cache, need_hidden):
past_len = cache.get_seq_length()
seq_len = inputs_embeds.shape[1]
cache_position = torch.arange(past_len, past_len + seq_len, device=inputs_embeds.device)
out = self.language_model(
inputs_embeds=inputs_embeds,
past_key_values=cache,
use_cache=True,
cache_position=cache_position,
position_ids=cache_position.unsqueeze(0),
output_hidden_states=need_hidden,
return_dict=True,
)
hidden = out.hidden_states[-1] if need_hidden else None
return out.logits, hidden, out.past_key_values
@torch.no_grad()
def _decode_children_batched(self, parent_cache, prompt_len, fork_indices,
eos_ids, max_new_tokens):
"""Decode all fork children concurrently as one batch (greedy AR).
Each child ``i`` inherits the parent prefix up to its ``<FORK>`` (real
length ``Li = prompt_len + fork_indices[i]``). The parent KV is copied
into a left-padded batched cache (row ``i`` right-aligned to hold its
``Li`` prefix positions), then all children step in lockstep until every
row hits EOS. Per-row ``position_ids`` carry each child's true positions
so RoPE matches the parent context. Returns a list of token-id lists.
Children run without the P-MTP head (batched ragged acceptance would
desync row lengths); the greedy output is identical to the serial path.
"""
device = self.device
B = len(fork_indices)
Ls = [prompt_len + f for f in fork_indices]
Lmax = max(Ls)
cache = DynamicCache()
for idx, layer in enumerate(parent_cache.layers):
K, V = layer.keys, layer.values # [1, h, Pseq, d]
h, d = K.shape[1], K.shape[3]
Kb = K.new_zeros((B, h, Lmax, d))
Vb = V.new_zeros((B, h, Lmax, d))
for i, Li in enumerate(Ls):
Kb[i, :, Lmax - Li:, :] = K[0, :, :Li, :]
Vb[i, :, Lmax - Li:, :] = V[0, :, :Li, :]
cache.update(Kb, Vb, idx)
mask = torch.zeros((B, Lmax), dtype=torch.long, device=device)
for i, Li in enumerate(Ls):
mask[i, Lmax - Li:] = 1
real_pos = torch.tensor(Ls, device=device) # <CHILD> sits at position Li
cur_embeds = self.language_model.get_input_embeddings()(
torch.tensor([[self.child_token_id]], device=device)).expand(B, 1, -1).contiguous()
outputs = [[] for _ in range(B)]
done = torch.zeros(B, dtype=torch.bool, device=device)
steps = 0
while steps < max_new_tokens and not bool(done.all()):
cache_len = cache.get_seq_length()
full_mask = torch.cat([mask, torch.ones((B, 1), dtype=torch.long, device=device)], dim=1)
out = self.language_model(
inputs_embeds=cur_embeds,
past_key_values=cache,
use_cache=True,
attention_mask=full_mask,
position_ids=real_pos.unsqueeze(1),
cache_position=torch.tensor([cache_len], device=device),
return_dict=True,
)
cache = out.past_key_values
nxt = out.logits[:, -1, :].argmax(-1) # [B]
mask = full_mask
real_pos = real_pos + 1
for i in range(B):
if not done[i]:
tid = int(nxt[i])
outputs[i].append(tid)
if tid in eos_ids or len(outputs[i]) >= max_new_tokens:
done[i] = True
cur_embeds = self.language_model.get_input_embeddings()(nxt.unsqueeze(1))
steps += 1
return outputs
@torch.no_grad()
def _decode_from(self, cache, first_embeds, eos_ids, max_new_tokens, use_mtp, k, allow_fork):
"""Greedily decode one branch (batch size 1) starting from ``cache``.
``first_embeds`` primes the branch (the prompt for the parent, or the
single ``<CHILD>`` token embedding for a child that inherits the parent
KV cache). With ``use_mtp`` the P-MTP head drafts ``k`` tokens per step
and the backbone verifies them in a single forward (greedy accept: keep
the longest matching prefix, then one correction/bonus token), rolling
the cache back with ``crop``. Returns ``(real_tokens, fork_indices,
cache)``; ``fork_indices`` marks positions of ``<FORK>`` tokens (only
when ``allow_fork``).
"""
device = first_embeds.device
logits, hidden, cache = self._lm_forward(first_embeds, cache, use_mtp)
q = cache.get_seq_length() - 1 # cache holds positions [0..q]
prod_hidden = hidden[:, -1, :] if use_mtp else None
t = int(logits[0, -1].argmax(-1)) # confirmed token at position q+1
real = []
fork_indices = []
def emit(tok):
real.append(tok)
if allow_fork and self.fork_token_id is not None and tok == self.fork_token_id:
fork_indices.append(len(real) - 1)
return tok in eos_ids
stopped = False
while len(real) < max_new_tokens and not stopped:
if use_mtp:
drafts = []
h = prod_hidden
prev = t
for _ in range(k):
prev_embeds = self.language_model.get_input_embeddings()(
torch.tensor([prev], device=device))
h = self.language_model.mtp.step(h, prev_embeds)
prev = int(self.language_model.lm_head(h).argmax(-1))
drafts.append(prev)
verify_ids = torch.tensor([[t] + drafts], device=device)
vlogits, vhidden, cache = self._lm_forward(self._embed_ids(verify_ids), cache, True)
target = vlogits[0].argmax(-1).tolist() # target[i] = argmax after verify_ids[i]
emitted = []
j = 0
all_match = True
for i in range(k):
emitted.append(target[i])
if drafts[i] != target[i]:
all_match = False
break
j += 1
if all_match:
emitted.append(target[k]) # bonus token
confirmed = [t] + emitted[:-1]
next_t = emitted[-1]
cache.crop((q + 1) + 1 + j) # keep [0..q+1+j]
prod_hidden = vhidden[0, j:j + 1, :]
q = q + 1 + j
t = next_t
for tok in confirmed:
if emit(tok):
stopped = True
break
else:
vlogits, _, cache = self._lm_forward(
self._embed_ids(torch.tensor([[t]], device=device)), cache, False)
if emit(t):
stopped = True
q = q + 1
t = int(vlogits[0, -1].argmax(-1))
if not stopped and len(real) < max_new_tokens:
emit(t)
return real, fork_indices, cache
@torch.no_grad()
def generate_hpd(self, tokenizer, pixel_values, question, generation_config=None,
use_mtp=False, num_speculative_tokens=6, num_patches_list=None,
IMG_START_TOKEN='<img>', IMG_END_TOKEN='</img>',
IMG_CONTEXT_TOKEN='<IMG_CONTEXT>', return_token_ids=False,
batch_children=False, verbose=False):
"""Hierarchical parallel decoding (``<FORK>``/``<CHILD>``) with optional P-MTP.
Reproduces the vLLM decoding paradigm in plain transformers for a single
image / single prompt (batch size 1):
- The parent "layout" branch is decoded greedily; each ``<FORK>`` token
spawns a content branch. The parent keeps the real ``<FORK>`` in its
context and continues.
- Each child branch inherits the parent KV cache (shared-prefix reuse,
like vLLM). By default children are decoded serially, reusing the
parent cache in place (crop to the tokens before each ``<FORK>``);
this is the fastest, lowest-memory option in transformers. Set
``batch_children=True`` to instead decode all children concurrently in
one left-padded batch (``_decode_children_batched``) -- identical
output, but for many short children over a long prefix it is slower
and much more memory-hungry (no paged prefix sharing like vLLM).
- Outputs are spliced: each ``<FORK>`` in the parent stream is rewritten
to ``<CHILD>`` and the matching child's tokens are inserted after it,
in fork order.
Use ``use_mtp=True`` only after ``load_mtp_weights(...)`` has been called.
"""
if self.fork_token_id is None or self.child_token_id is None:
raise RuntimeError('fork_token_id/child_token_id missing from config; cannot fork-decode.')
if use_mtp and not self._mtp_weights_loaded:
raise RuntimeError('P-MTP weights not loaded. Call model.load_mtp_weights(path) '
'or set use_mtp=False.')
generation_config = generation_config or {}
max_new_tokens = generation_config.get('max_new_tokens', 8000)
k = num_speculative_tokens
if pixel_values is not None and '<image>' not in question:
question = '<image>\n' + question
if num_patches_list is None:
num_patches_list = [pixel_values.shape[0]] if pixel_values is not None else []
self.img_context_token_id = tokenizer.convert_tokens_to_ids(IMG_CONTEXT_TOKEN)
template = get_conv_template(self.template)
template.system_message = self.system_message
template.append_message(template.roles[0], question)
template.append_message(template.roles[1], None)
query = template.get_prompt()
for num_patches in num_patches_list:
image_tokens = IMG_START_TOKEN + IMG_CONTEXT_TOKEN * self.num_image_token * num_patches + IMG_END_TOKEN
query = query.replace('<image>', image_tokens, 1)
prompt_ids = tokenizer(query, return_tensors='pt')['input_ids'].to(self.device)
prompt_len = prompt_ids.shape[1]
eos_ids = set()
sep_id = tokenizer.convert_tokens_to_ids(template.sep.strip())
if isinstance(sep_id, int) and sep_id >= 0:
eos_ids.add(sep_id)
if self.config.eos_token_id is not None:
eos_ids.add(self.config.eos_token_id)
vit_embeds = self.extract_feature(pixel_values) if pixel_values is not None else None
prompt_embeds = self._embed_ids(prompt_ids, vit_embeds)
parent_real, fork_indices, parent_cache = self._decode_from(
DynamicCache(), prompt_embeds, eos_ids, max_new_tokens, use_mtp, k, allow_fork=True)
# Once the parent finishes, every fork point and its shared prefix is
# fixed. Children can then be decoded either serially (reusing the parent
# cache in place, low memory) or concurrently in one padded batch
# (``batch_children=True``). NOTE: for many short children over a long
# shared prefix, batching is slower and far more memory-hungry in stock
# transformers (the prefix KV is duplicated per row and short children
# still step until the longest one finishes), so serial is the default.
children = []
if fork_indices and batch_children:
children = self._decode_children_batched(
parent_cache, prompt_len, fork_indices, eos_ids, max_new_tokens)
elif fork_indices:
child_embeds = self.language_model.get_input_embeddings()(
torch.tensor([[self.child_token_id]], device=self.device))
children = [None] * len(fork_indices)
# Reverse fork order: crop the shared parent cache to the tokens
# before each <FORK> and reuse it in place. A child only appends
# beyond its crop point, which the next (smaller) crop discards, so
# the parent prefix KV stays intact and no clone is needed.
for i in range(len(fork_indices) - 1, -1, -1):
parent_cache.crop(prompt_len + fork_indices[i])
child_real, _, _ = self._decode_from(
parent_cache, child_embeds, eos_ids, max_new_tokens, use_mtp, k, allow_fork=False)
children[i] = child_real
final_ids = []
child_ptr = 0
for tok in parent_real:
if tok == self.fork_token_id:
final_ids.append(self.child_token_id)
if child_ptr < len(children):
final_ids.extend(children[child_ptr])
child_ptr += 1
else:
final_ids.append(tok)
text = tokenizer.decode(final_ids, skip_special_tokens=True)
if verbose:
print(f'[HPD] parent tokens={len(parent_real)}, forks={len(fork_indices)}, '
f'children_tokens={[len(c) for c in children]}')
if return_token_ids:
return text, final_ids
return text
@property
def lm_head(self):
return self.language_model.get_output_embeddings()
def get_output_embeddings(self):
return self.language_model.get_output_embeddings()
def get_input_embeddings(self):
return self.language_model.get_input_embeddings()
def set_input_embeddings(self, value):
return self.language_model.set_input_embeddings(value)
def set_output_embeddings(self, value):
return self.language_model.set_output_embeddings(value)