From 977a5543538796f7b2c84904eddb45fff574378b Mon Sep 17 00:00:00 2001 From: grimoire Date: Thu, 23 Jul 2026 21:22:11 +0800 Subject: [PATCH 01/16] Add initial ConceptLM VQ PyTorch modules --- lmdeploy/pytorch/configurations/conceptlm.py | 59 + lmdeploy/pytorch/models/intern_ncp.py | 1394 ++++++++++++++++++ lmdeploy/pytorch/models/module_map.py | 5 + 3 files changed, 1458 insertions(+) create mode 100644 lmdeploy/pytorch/configurations/conceptlm.py create mode 100644 lmdeploy/pytorch/models/intern_ncp.py diff --git a/lmdeploy/pytorch/configurations/conceptlm.py b/lmdeploy/pytorch/configurations/conceptlm.py new file mode 100644 index 0000000000..b16210249b --- /dev/null +++ b/lmdeploy/pytorch/configurations/conceptlm.py @@ -0,0 +1,59 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from lmdeploy.utils import get_logger + +from .builder import AutoModelConfigBuilder +from .default import DefaultModelConfigBuilder + +logger = get_logger('lmdeploy') + + +class ConceptLMModelConfigBuilder(AutoModelConfigBuilder): + """Config builder for ConceptLM V2.2-VQ. + + The upstream checkpoint config does not declare bos/eos/pad token ids, + so derive them from the tokenizer when missing before handing the config + to the default builder (which asserts they exist). + """ + + @classmethod + def condition(cls, hf_config): + """config.""" + archs = getattr(hf_config, 'architectures', None) or [] + return 'ConceptLMV22VQForCausalLM' in archs + + @classmethod + def build(cls, hf_config, model_path: str = None, **kwargs): + """build.""" + # fill missing special token ids from the tokenizer + if getattr(hf_config, 'bos_token_id', None) is None: + cls._fill_special_tokens(hf_config, model_path) + + model_config = DefaultModelConfigBuilder.build(hf_config, model_path, **kwargs) + + enc_layers = int(hf_config.concept_encoder_layers) + concept_layers = int(hf_config.concept_special_layers) + dec_layers = int(hf_config.concept_decoder_layers) + model_config.num_layers = enc_layers + concept_layers + dec_layers + + # TODO: ConceptLM's concept predictor attends over a compressed + # chunk-level timeline, so these concept KV layers do not need full + # token-length block capacity. Keep the standard KV cache layout for + # now to avoid adding another cache-engine abstraction; optimize this + # later if concept KV memory becomes a real bottleneck. + model_config.llm_config.concept_kv_encoder_offset = 0 + model_config.llm_config.concept_kv_concept_offset = enc_layers + model_config.llm_config.concept_kv_decoder_offset = enc_layers + concept_layers + model_config.llm_config.concept_kv_total_layers = model_config.num_layers + return model_config + + @staticmethod + def _fill_special_tokens(hf_config, model_path: str = None): + try: + from transformers import AutoTokenizer + tok = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True) + hf_config.bos_token_id = tok.bos_token_id + hf_config.eos_token_id = tok.eos_token_id + if getattr(hf_config, 'pad_token_id', None) is None: + hf_config.pad_token_id = tok.pad_token_id + except Exception as e: # noqa: BLE001 + logger.warning(f'ConceptLM: failed to derive special token ids: {e}') diff --git a/lmdeploy/pytorch/models/intern_ncp.py b/lmdeploy/pytorch/models/intern_ncp.py new file mode 100644 index 0000000000..9feaa17ad1 --- /dev/null +++ b/lmdeploy/pytorch/models/intern_ncp.py @@ -0,0 +1,1394 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""lmdeploy adapter for ConceptLM V2.2-VQ. + +Reference implementation: +``concept_olmo_stage_2_V1/modeling_conceptlm_v22_vq.py`` + +Modules are added incrementally. Current state: + - token embedding + output projection (lm_head) + - ``_OlmoBlock`` (encoder/decoder/concept_predictor backbone): attention, + mlp, rmsnorm, rope. Wired to lmdeploy primitives so it is TP-correct and + ready to plug into the engine's paged attention path. + - ``_Quantizer``: stacked VQ codebook parameter, replicated across TP. + - ``_SelfDD``: replicated per-token depth mixer for encoder hidden history. + - ``_ResidualRoute``: replicated residual source mixer for decoder routes. + - ``_TwoRouteAdd``: decoder depth mixing plus final-concept residual route. + - ``_ConceptPredictor``: concept block container and prediction heads. + +Results are NOT yet correct end-to-end — the encoder/concept/decoder control +flow is added in later steps. The module structure mirrors the reference for +readability. +""" + +from collections.abc import Iterable +from typing import Any + +import torch +from torch import nn +from torch.nn import functional as F +from transformers.configuration_utils import PretrainedConfig + +from lmdeploy.pytorch.model_inputs import StepContext, StepContextManager +from lmdeploy.pytorch.nn import ApplyRotaryEmb, Attention, RMSNorm, SiluAndMul, build_rotary_embedding +from lmdeploy.pytorch.nn.linear import (build_down_linear, build_gateup_linear, + build_merged_colwise_linear, build_o_proj, build_qkv_proj) +from lmdeploy.pytorch.weight_loader.model_weight_loader import load_weight + +from .patch import add_prefix +from .utils.cudagraph import CudaGraphMixin +from .utils.model import DeployModelMixinV1, build_embedding + +_CONFIG_VALUE = object() +_HistoryStates = list[torch.Tensor] | tuple[torch.Tensor, ...] | torch.Tensor +_SourceStates = list[torch.Tensor] | tuple[torch.Tensor, ...] | torch.Tensor | None + + +def _get_configured_window(config: PretrainedConfig): + """Return the reference OLMo window setting as an int or None.""" + window_size = getattr(config, 'window_size', None) + if window_size is None: + return None + if isinstance(window_size, (list, tuple)): + window_size = window_size[0] + if window_size is None: + return None + window_size = int(window_size) + return window_size if window_size > 0 else None + + +def _make_olmo_rotary_embedding(config: PretrainedConfig, + device: torch.device = None) -> nn.Module: + """Build ConceptLM/OLMo rotary embedding from its native config fields.""" + rotary_interleaved = bool(getattr(config, 'rotary_interleaved', False)) + if rotary_interleaved: + raise NotImplementedError('ConceptLM rotary_interleaved=True is not supported by the LMDeploy block yet.') + + head_dim = int(config.kv_channels) + rotary_percent = float(getattr(config, 'rotary_percent', 1.0)) + rotary_dim = int(head_dim * rotary_percent) + rotary_dim -= rotary_dim % 2 + if rotary_dim <= 0: + raise ValueError(f'Invalid ConceptLM rotary dimension: head_dim={head_dim}, rotary_percent={rotary_percent}') + + partial_rotary_factor = rotary_dim / head_dim + return build_rotary_embedding( + dim=head_dim, + max_position_embeddings=getattr(config, 'max_position_embeddings', getattr(config, 'max_sequence_length', + 2048)), + base=getattr(config, 'rotary_base', 10000), + partial_rotary_factor=partial_rotary_factor, + device=device, + ) + + +def _repack_olmo_qkv_weight(loaded_weight: torch.Tensor, num_heads: int, head_dim: int): + """Convert native OLMo per-head [Q,K,V] QKV packing to LMDeploy [Q][K][V].""" + leading_shape = loaded_weight.shape[1:] + loaded_weight = loaded_weight.reshape(num_heads, 3, head_dim, *leading_shape) + query = loaded_weight[:, 0].flatten(0, 1) + key = loaded_weight[:, 1].flatten(0, 1) + value = loaded_weight[:, 2].flatten(0, 1) + return torch.cat([query, key, value], dim=0) + + +def _load_stacked_codebook_weight(param: torch.nn.Parameter, loaded_weight: torch.Tensor, codebook_idx: int): + """Load one native ``codebook.N`` checkpoint tensor into a stacked codebook.""" + assert 0 <= codebook_idx < param.size(0), f'Invalid codebook index: {codebook_idx}' + target = param.data[codebook_idx] + assert target.size() == loaded_weight.size(), ( + f'Attempted to load codebook weight ({loaded_weight.size()}) into parameter slice ({target.size()})') + target.copy_(loaded_weight) + + +def _qk_rmsnorm_variance(query: torch.Tensor, key: torch.Tensor) -> torch.Tensor: + """Local Q/K squared sums before TP all-reduce.""" + query = query.float() + key = key.float() + query_var = (query * query).sum(-1, keepdim=True) + key_var = (key * key).sum(-1, keepdim=True) + return torch.stack([query_var, key_var], dim=0) + + +def _qk_rmsnorm_apply(query: torch.Tensor, + key: torch.Tensor, + variance: torch.Tensor, + query_weight: torch.Tensor, + key_weight: torch.Tensor, + hidden_size: int, + eps: float): + """Apply whole-hidden Q/K RMSNorm from an already all-reduced variance.""" + dtype = query.dtype + query_var, key_var = variance / hidden_size + eps + query = (query.float() * torch.rsqrt(query_var)).to(dtype) * query_weight + key = (key.float() * torch.rsqrt(key_var)).to(dtype) * key_weight + return query, key + + +class ConceptLMV22VQEmbedding(nn.Module): + """Token embedding container. + + Mirrors ``self.embedding`` in the reference: a plain ``nn.Module`` holding + a ``word_embeddings`` table. Keeping the attribute name ``word_embeddings`` + makes the checkpoint key ``embedding.word_embeddings.weight`` line up + directly with ``self.embedding.word_embeddings``. + """ + + def __init__(self, + config: PretrainedConfig, + dtype: torch.dtype = None, + device: torch.device = None): + super().__init__() + self.word_embeddings = build_embedding( + config.vocab_size, + config.hidden_size, + getattr(config, 'pad_token_id', None), + dtype=dtype, + device=device, + ) + + def forward(self, input_ids: torch.Tensor) -> torch.Tensor: + """forward.""" + return self.word_embeddings(input_ids) + + +class ConceptLMV22VQQuantizer(nn.Module): + """Rewrite of ``_Quantizer``. + + The reference keeps codebooks as a ``ParameterList`` only to produce + checkpoint keys ``codebook.0`` ... ``codebook.N``. Runtime computation only + needs the stacked tensor returned by ``transformed_codebook()``. + """ + + def __init__(self, + config: PretrainedConfig, + dtype: torch.dtype = None, + device: torch.device = None): + super().__init__() + self.num_codebooks = int(config.concept_v22_vq_num_codebooks) + self.codebook_size = int(config.concept_v22_vq_codebook_size) + hidden_size = int(config.hidden_size) + assert hidden_size % self.num_codebooks == 0, ( + f'hidden_size={hidden_size} must be divisible by num_codebooks={self.num_codebooks}') + self.codebook_dim = hidden_size // self.num_codebooks + self.hidden_size = hidden_size + self.codebook = nn.Parameter( + torch.empty( + self.num_codebooks, + self.codebook_size, + self.codebook_dim, + dtype=dtype, + device=device, + ), + requires_grad=False, + ) + + def transformed_codebook(self): + """Return codebook as ``[num_codebooks, codebook_size, codebook_dim]``.""" + return self.codebook + + def forward(self, concept_logits: torch.Tensor) -> torch.Tensor: + """Convert per-codebook logits to hidden vectors. + + Args: + concept_logits: ``[..., num_codebooks, codebook_size]``. In the + LMDeploy engine the leading dims may be a packed continuous + batching dimension, e.g. ``[num_concepts_total]``. + + Returns: + Quantized/predicted vectors with shape ``[..., hidden_size]``. + """ + assert concept_logits.shape[-2:] == (self.num_codebooks, self.codebook_size), ( + f'Expected concept logits trailing shape {(self.num_codebooks, self.codebook_size)}, ' + f'got {tuple(concept_logits.shape[-2:])}.') + codebook = self.transformed_codebook().to(concept_logits.dtype) + vectors = torch.einsum('...hk,hkd->...hd', concept_logits, codebook) + return vectors.flatten(-2, -1) + + +class ConceptLMV22VQDepthDD(nn.Module): + """Rewrite of ``_DepthDD``. + + This is a small replicated per-token depth mixer. It does not mix tokens; + it computes ``num_prev`` route weights from the current hidden state and + combines the matching per-layer history states. The reference only handles + dense ``[seq, batch, hidden]`` tensors because it stacks history at + ``dim=2``. LMDeploy continuous batching commonly uses packed + ``[num_tokens, hidden]`` tensors. Runtime should pass a preallocated tensor + history to avoid repeated ``torch.stack`` copies; list/tuple input remains + only as a debug/parity convenience. + """ + + def __init__(self, + config: PretrainedConfig, + layer_idx: int, + use_softmax: bool, + dtype: torch.dtype = None, + device: torch.device = None): + super().__init__() + self.num_prev = int(layer_idx) + 2 + route_hidden_size = self.num_prev + self.w1 = nn.Linear(config.hidden_size, route_hidden_size, bias=False, dtype=dtype, device=device) + self.w2 = nn.Linear(route_hidden_size, self.num_prev, bias=False, dtype=dtype, device=device) + self.static_a = nn.Parameter(torch.zeros(self.num_prev, dtype=dtype, device=device), requires_grad=False) + self.use_softmax = bool(use_softmax) + for param in self.parameters(): + param.requires_grad_(False) + + def _history_tensor(self, hidden_states: torch.Tensor, history_states: _HistoryStates, history_dim: int): + """Return history in shape ``[..., num_prev, hidden_size]``.""" + if not isinstance(history_states, torch.Tensor): + assert len(history_states) == self.num_prev, ( + f'Expected {self.num_prev} history states, got {len(history_states)}.') + history_states = torch.stack(tuple(history_states), dim=-2) + history_dim = -2 + + history_dim = history_dim if history_dim >= 0 else history_dim + history_states.dim() + assert 0 <= history_dim < history_states.dim(), f'Invalid history_dim={history_dim}.' + assert history_states.shape[history_dim] == self.num_prev, ( + f'Expected history dimension {history_dim} to be {self.num_prev}, ' + f'got {history_states.shape[history_dim]}.') + if history_dim != history_states.dim() - 2: + history_states = history_states.movedim(history_dim, -2) + assert history_states.shape[:-2] == hidden_states.shape[:-1], ( + f'Expected history leading shape {tuple(hidden_states.shape[:-1])}, ' + f'got {tuple(history_states.shape[:-2])}.') + assert history_states.shape[-1] == hidden_states.shape[-1], ( + f'Expected history hidden size {hidden_states.shape[-1]}, got {history_states.shape[-1]}.') + return history_states + + def forward(self, + hidden_states: torch.Tensor, + history_states: _HistoryStates, + history_dim: int = -2) -> torch.Tensor: + """forward.""" + history = self._history_tensor(hidden_states, history_states, history_dim) + weights = self.w2(F.gelu(self.w1(hidden_states))) + weights = weights + self.static_a.to(dtype=weights.dtype) + if self.use_softmax: + weights = weights.softmax(dim=-1) + return torch.einsum('...l,...lh->...h', weights, history) + + +class ConceptLMV22VQSelfDD(nn.Module): + """Rewrite of ``_SelfDD``.""" + + def __init__(self, + config: PretrainedConfig, + num_layers: int, + use_softmax: bool = False, + dtype: torch.dtype = None, + device: torch.device = None): + super().__init__() + self.num_layers = int(num_layers) + self.depth_dds = nn.ModuleList([ + ConceptLMV22VQDepthDD(config, layer_idx, use_softmax, dtype=dtype, device=device) + for layer_idx in range(self.num_layers) + ]) + + def make_history_buffer(self, hidden_states: torch.Tensor) -> torch.Tensor: + """Allocate layer-major history buffer ``[num_layers + 1, *hidden_shape]``.""" + return hidden_states.new_empty((self.num_layers + 1, *hidden_states.shape)) + + @staticmethod + def write_history(history_buffer: torch.Tensor, slot_idx: int, hidden_states: torch.Tensor): + """Copy one history block into a layer-major history buffer.""" + assert history_buffer.dim() == hidden_states.dim() + 1, ( + f'Expected history buffer dim {hidden_states.dim() + 1}, got {history_buffer.dim()}.') + assert history_buffer.shape[1:] == hidden_states.shape, ( + f'Expected history buffer trailing shape {tuple(hidden_states.shape)}, ' + f'got {tuple(history_buffer.shape[1:])}.') + history_buffer[int(slot_idx)].copy_(hidden_states) + return history_buffer + + @staticmethod + def history_view(history_buffer: torch.Tensor, layer_idx: int): + """Return layer-major history needed by ``layer_idx`` without copying. + + This is CUDA-graph safe when ``layer_idx`` is a Python constant for the + current layer and ``history_buffer`` has the fixed graph-capture shape. + """ + return history_buffer[:int(layer_idx) + 2] + + def forward_from_buffer(self, + layer_idx: int, + hidden_states: torch.Tensor, + history_buffer: torch.Tensor): + """Runtime path: fixed buffer, no list materialization.""" + layer_idx = int(layer_idx) + return self.depth_dds[layer_idx]( + hidden_states, + self.history_view(history_buffer, layer_idx), + history_dim=0, + ) + + def forward(self, + layer_idx: int, + hidden_states: torch.Tensor, + history_states: _HistoryStates): + """forward.""" + return self.depth_dds[int(layer_idx)](hidden_states, history_states) + + +class ConceptLMV22VQResidualRoute(nn.Module): + """Rewrite of ``_ResidualRoute``. + + This module computes a source-state mixture and adds it as a gated residual + update to the target hidden state. It is small and replicated. + + Runtime/static-graph code should use ``forward_from_buffer`` with a full + fixed-size source buffer. Flexible active-source/list paths are retained for + parity tests and WIP reference wiring, but they should not be used inside a + captured CUDA graph because they can change intermediate shapes. + """ + + def __init__(self, + config: PretrainedConfig, + num_source_states: int, + use_softmax: bool = True, + dtype: torch.dtype = None, + device: torch.device = None): + super().__init__() + self.num_source_states = int(num_source_states) + route_hidden_size = max(1, self.num_source_states) + self.w1 = nn.Linear(config.hidden_size, route_hidden_size, bias=False, dtype=dtype, device=device) + self.w2 = nn.Linear(route_hidden_size, self.num_source_states, bias=False, dtype=dtype, device=device) + self.residual_diag = nn.Parameter(torch.zeros(config.hidden_size, dtype=dtype, device=device), + requires_grad=False) + self.use_softmax = bool(use_softmax) + for param in self.parameters(): + param.requires_grad_(False) + + def _source_tensor(self, + target_hidden: torch.Tensor, + source_states: _SourceStates, + source_dim: int, + expected_leading_shape: tuple[int, ...] | None = None): + """Return source states in shape ``[..., active_sources, hidden_size]``.""" + if source_states is None: + return None + if not isinstance(source_states, torch.Tensor): + if len(source_states) == 0: + return None + source_states = torch.stack(tuple(source_states), dim=-2) + source_dim = -2 + + source_dim = source_dim if source_dim >= 0 else source_dim + source_states.dim() + assert 0 <= source_dim < source_states.dim(), f'Invalid source_dim={source_dim}.' + if source_states.shape[source_dim] == 0: + return None + if source_dim != source_states.dim() - 2: + source_states = source_states.movedim(source_dim, -2) + if expected_leading_shape is None: + expected_leading_shape = target_hidden.shape[:-1] + assert source_states.shape[:-2] == expected_leading_shape, ( + f'Expected source leading shape {tuple(expected_leading_shape)}, ' + f'got {tuple(source_states.shape[:-2])}.') + assert source_states.shape[-1] == target_hidden.shape[-1], ( + f'Expected source hidden size {target_hidden.shape[-1]}, got {source_states.shape[-1]}.') + return source_states + + def _route_weights(self, target_hidden: torch.Tensor, active_sources: int): + """Compute route weights and keep the last active source logits.""" + weights = self.w2(F.gelu(self.w1(target_hidden))) + weights = weights[..., -active_sources:] + if self.use_softmax: + weights = weights.softmax(dim=-1) + return weights + + def _add_update(self, + target_hidden: torch.Tensor, + source_mix: torch.Tensor, + residual_scale: torch.Tensor | None = None): + """Apply residual diagonal and optional scale, then add to target.""" + update = source_mix * self.residual_diag.to(source_mix.dtype) + if residual_scale is not None: + update = update * residual_scale.to(update.dtype) + return target_hidden + update.to(target_hidden.dtype) + + def make_source_buffer(self, hidden_states: torch.Tensor) -> torch.Tensor: + """Allocate source-major buffer ``[num_source_states, *hidden_shape]``.""" + return hidden_states.new_empty((self.num_source_states, *hidden_states.shape)) + + @staticmethod + def write_source(source_buffer: torch.Tensor, slot_idx: int, source_state: torch.Tensor): + """Copy one source state into a source-major buffer.""" + assert source_buffer.dim() == source_state.dim() + 1, ( + f'Expected source buffer dim {source_state.dim() + 1}, got {source_buffer.dim()}.') + assert source_buffer.shape[1:] == source_state.shape, ( + f'Expected source buffer trailing shape {tuple(source_state.shape)}, ' + f'got {tuple(source_buffer.shape[1:])}.') + source_buffer[int(slot_idx)].copy_(source_state) + return source_buffer + + @staticmethod + def source_view(source_buffer: torch.Tensor, active_sources: int | None = None): + """Return active source-major view without copying. + + Passing ``active_sources`` is a flexible/debug path. Static graph + runtime should pass full fixed buffers and leave ``active_sources`` as + ``None``. + """ + if active_sources is None: + return source_buffer + return source_buffer[:int(active_sources)] + + def forward(self, + target_hidden: torch.Tensor, + source_states: _SourceStates, + residual_scale: torch.Tensor | None = None, + source_dim: int = -2): + """Flexible/debug forward path. + + This accepts lists, ``None``, and active source tensors for reference + parity. Use ``forward_from_buffer`` for static-shape runtime. + """ + source_states = self._source_tensor(target_hidden, source_states, source_dim) + if source_states is None: + return target_hidden + active_sources = source_states.shape[-2] + weights = self._route_weights(target_hidden, active_sources) + source_mix = torch.einsum('...m,...mh->...h', weights, source_states) + return self._add_update(target_hidden, source_mix, residual_scale) + + def forward_from_buffer(self, + target_hidden: torch.Tensor, + source_buffer: torch.Tensor, + residual_scale: torch.Tensor | None = None): + """Runtime path: read a full fixed source-major buffer. + + ``source_buffer`` shape is ``[num_source_states, *target_hidden.shape]``. + This keeps source count fixed across CUDA graph capture/replay. + """ + assert source_buffer.shape[0] == self.num_source_states, ( + f'Expected full source buffer with {self.num_source_states} states, got {source_buffer.shape[0]}.') + return self.forward( + target_hidden, + source_buffer, + residual_scale=residual_scale, + source_dim=0, + ) + + def forward_active_from_buffer(self, + target_hidden: torch.Tensor, + source_buffer: torch.Tensor, + active_sources: int, + residual_scale: torch.Tensor | None = None): + """Flexible/debug buffer path with reduced active source count.""" + return self.forward( + target_hidden, + self.source_view(source_buffer, active_sources), + residual_scale=residual_scale, + source_dim=0, + ) + + def forward_repeated_chunks(self, + target_hidden: torch.Tensor, + source_states: _SourceStates, + chunk_size: int, + shift_feature: bool, + residual_scale: torch.Tensor | None = None, + source_dim: int = 2): + """Reference repeated-chunk route used by decoder-read-concept. + + ``target_hidden`` is dense ``[seq, batch, hidden]`` and source states are + chunk-level ``[chunks, batch, sources, hidden]`` after ``source_dim`` is + normalized to 2. A packed continuous-batching runtime will need token to + chunk metadata before using this path end-to-end. This method is not a + complete CUDA-graph runtime path yet because chunk lengths still need a + fixed-buffer/mask contract at the caller level. + """ + if source_states is None: + return target_hidden + assert target_hidden.dim() == 3, ( + f'forward_repeated_chunks expects [seq, batch, hidden], got {tuple(target_hidden.shape)}.') + + if not isinstance(source_states, torch.Tensor): + if len(source_states) == 0: + return target_hidden + source_states = torch.stack(tuple(source_states), dim=2) + source_dim = 2 + source_dim = source_dim if source_dim >= 0 else source_dim + source_states.dim() + assert 0 <= source_dim < source_states.dim(), f'Invalid source_dim={source_dim}.' + if source_states.shape[source_dim] == 0: + return target_hidden + if source_dim != 2: + source_states = source_states.movedim(source_dim, 2) + + seq_len, batch_size, hidden_size = target_hidden.shape + assert source_states.dim() == 4, ( + f'Expected source states [chunks, batch, sources, hidden], got {tuple(source_states.shape)}.') + assert source_states.shape[1] == batch_size, ( + f'Expected source batch size {batch_size}, got {source_states.shape[1]}.') + assert source_states.shape[3] == hidden_size, ( + f'Expected source hidden size {hidden_size}, got {source_states.shape[3]}.') + num_chunks, _, active_sources, _ = source_states.shape + chunk_size = int(chunk_size) + weights = self._route_weights(target_hidden, active_sources) + if shift_feature: + weights = torch.cat((weights.new_zeros(1, batch_size, active_sources), weights), dim=0) + repeated_len = num_chunks * chunk_size + if weights.shape[0] < repeated_len: + pad_len = repeated_len - weights.shape[0] + weights = torch.cat((weights, weights.new_zeros(pad_len, batch_size, active_sources)), dim=0) + weights = weights[:repeated_len] + + source_mix = torch.einsum( + 'ckbm,cbmh->ckbh', + weights.reshape(num_chunks, chunk_size, batch_size, active_sources), + source_states, + ).reshape(repeated_len, batch_size, hidden_size) + if shift_feature: + source_mix = source_mix[1:1 + seq_len] + else: + source_mix = source_mix[:seq_len] + if source_mix.shape[0] < seq_len: + pad_len = seq_len - source_mix.shape[0] + source_mix = torch.cat((source_mix, source_mix.new_zeros(pad_len, batch_size, hidden_size)), dim=0) + return self._add_update(target_hidden, source_mix, residual_scale) + + def forward_repeated_chunks_from_buffer(self, + target_hidden: torch.Tensor, + source_buffer: torch.Tensor, + chunk_size: int, + shift_feature: bool, + residual_scale: torch.Tensor | None = None): + """Full source-major buffer variant of ``forward_repeated_chunks``.""" + assert source_buffer.shape[0] == self.num_source_states, ( + f'Expected full source buffer with {self.num_source_states} states, got {source_buffer.shape[0]}.') + return self.forward_repeated_chunks( + target_hidden, + source_buffer, + chunk_size, + shift_feature, + residual_scale=residual_scale, + source_dim=0, + ) + + def forward_repeated_chunks_active_from_buffer(self, + target_hidden: torch.Tensor, + source_buffer: torch.Tensor, + chunk_size: int, + shift_feature: bool, + active_sources: int, + residual_scale: torch.Tensor | None = None): + """Flexible/debug repeated-chunk buffer path with reduced active source count.""" + return self.forward_repeated_chunks( + target_hidden, + self.source_view(source_buffer, active_sources), + chunk_size, + shift_feature, + residual_scale=residual_scale, + source_dim=0, + ) + + +class ConceptLMV22VQConceptRoute(nn.Module): + """Rewrite of ``_ConceptRoute``. + + Applies LayerNorm to the final concept state, scales it elementwise with a + learned diagonal, optionally applies a route scale, then adds the update to + decoder hidden states. This is replicated and token-local. + """ + + def __init__(self, + config: PretrainedConfig, + dtype: torch.dtype = None, + device: torch.device = None): + super().__init__() + self.concept_norm = nn.LayerNorm(config.hidden_size, + eps=getattr(config, 'layernorm_epsilon', 1e-6), + dtype=dtype, + device=device) + self.final_diag = nn.Parameter(torch.zeros(config.hidden_size, dtype=dtype, device=device), + requires_grad=False) + for param in self.parameters(): + param.requires_grad_(False) + + def forward(self, + hidden_states: torch.Tensor, + final_concept_state: torch.Tensor, + final_scale: torch.Tensor | None = None): + """forward.""" + concept = self.concept_norm(final_concept_state) + update = concept * self.final_diag.to(concept.dtype) + if final_scale is not None: + update = update * final_scale.to(update.dtype) + return hidden_states + update.to(hidden_states.dtype) + + +class ConceptLMV22VQTwoRouteAdd(nn.Module): + """Rewrite of ``_TwoRouteAdd``. + + It first applies decoder-side ``_DepthDD`` to the decoder history, then + injects the final concept state with ``_ConceptRoute``. + """ + + def __init__(self, + config: PretrainedConfig, + dtype: torch.dtype = None, + device: torch.device = None): + super().__init__() + self.num_layers = int(config.concept_decoder_layers) + use_softmax = bool(getattr(config, 'concept_dd_two_route_add_decoder_use_softmax', True)) + self.decoder_dds = nn.ModuleList([ + ConceptLMV22VQDepthDD(config, layer_idx, use_softmax, dtype=dtype, device=device) + for layer_idx in range(self.num_layers) + ]) + self.concept_routes = nn.ModuleList([ + ConceptLMV22VQConceptRoute(config, dtype=dtype, device=device) + for _ in range(self.num_layers) + ]) + + def make_history_buffer(self, hidden_states: torch.Tensor) -> torch.Tensor: + """Allocate layer-major decoder history buffer ``[num_layers + 1, *hidden_shape]``.""" + return hidden_states.new_empty((self.num_layers + 1, *hidden_states.shape)) + + @staticmethod + def write_history(history_buffer: torch.Tensor, slot_idx: int, hidden_states: torch.Tensor): + """Copy one decoder history block into a layer-major history buffer.""" + return ConceptLMV22VQSelfDD.write_history(history_buffer, slot_idx, hidden_states) + + @staticmethod + def history_view(history_buffer: torch.Tensor, layer_idx: int): + """Return layer-major decoder history needed by ``layer_idx`` without copying.""" + return ConceptLMV22VQSelfDD.history_view(history_buffer, layer_idx) + + def forward_from_buffer(self, + layer_idx: int, + hidden_states: torch.Tensor, + history_buffer: torch.Tensor, + final_concept_state: torch.Tensor, + final_scale: torch.Tensor | None = None): + """Runtime path: fixed decoder history buffer, no list materialization.""" + layer_idx = int(layer_idx) + hidden_states = self.decoder_dds[layer_idx]( + hidden_states, + self.history_view(history_buffer, layer_idx), + history_dim=0, + ) + return self.concept_routes[layer_idx](hidden_states, final_concept_state, final_scale) + + def forward(self, + layer_idx: int, + hidden_states: torch.Tensor, + history_states: _HistoryStates, + final_concept_state: torch.Tensor, + final_scale: torch.Tensor | None = None): + """Flexible/debug forward path.""" + layer_idx = int(layer_idx) + hidden_states = self.decoder_dds[layer_idx](hidden_states, history_states) + return self.concept_routes[layer_idx](hidden_states, final_concept_state, final_scale) + + +class ConceptLMV22VQPredictionHeads(nn.Module): + """Merged per-codebook prediction heads. + + Native checkpoints store one ``prediction_heads.N`` linear per concept + codebook. Runtime uses one merged projection and loads each native head into + a deterministic output shard through LMDeploy's ``param.weight_loader`` + pattern. + """ + + def __init__(self, + config: PretrainedConfig, + dtype: torch.dtype = None, + device: torch.device = None, + prefix: str = ''): + codebook_size = int(config.concept_v22_vq_codebook_size) + num_codebooks = int(config.concept_v22_vq_num_codebooks) + super().__init__() + self.num_codebooks = num_codebooks + self.codebook_size = codebook_size + quantization_config = getattr(config, 'quantization_config', None) + self.proj = build_merged_colwise_linear( + config.hidden_size, + [codebook_size] * num_codebooks, + bias=True, + dtype=dtype, + device=device, + quant_config=quantization_config, + # Keep concept logits replicated for now. Output TP would make the + # trailing ``[num_codebooks, codebook_size]`` contract sharded and + # needs a defined distributed sampling/gather path first. + is_tp=False, + out_names=list(range(num_codebooks)), + prefix=add_prefix('proj', prefix), + ) + + def forward(self, hidden_states: torch.Tensor): + """Return logits in shape ``[..., num_codebooks, codebook_size]``.""" + logits = self.proj(hidden_states) + return logits.unflatten(-1, (self.num_codebooks, self.codebook_size)) + + +class ConceptLMV22VQConceptPredictor(nn.Module): + """Rewrite of ``_ConceptPredictor``. + + The predictor owns the high-level concept OLMo block, per-codebook + prediction heads, concept self-DD, encoder-read routes, and the shared + source LayerNorm used before encoder states are routed into the concept + stream. + """ + + def __init__(self, + config: PretrainedConfig, + dtype: torch.dtype = None, + device: torch.device = None, + prefix: str = ''): + super().__init__() + self.num_layers = int(config.concept_special_layers) + self.num_codebooks = int(config.concept_v22_vq_num_codebooks) + self.codebook_size = int(config.concept_v22_vq_codebook_size) + self.hlm_block = ConceptLMV22VQOlmoBlock(config, + self.num_layers, + post_layer_norm=True, + dtype=dtype, + device=device, + prefix=add_prefix('hlm_block', prefix)) + self.prediction_heads = ConceptLMV22VQPredictionHeads(config, + dtype=dtype, + device=device, + prefix=add_prefix('prediction_heads', prefix)) + self.concept_self_dd = ConceptLMV22VQSelfDD(config, + self.num_layers, + use_softmax=False, + dtype=dtype, + device=device) + self.concept_read_encoder_routes = nn.ModuleList([ + ConceptLMV22VQResidualRoute(config, + int(config.concept_encoder_layers) - 1, + use_softmax=True, + dtype=dtype, + device=device) + for _ in range(self.num_layers) + ]) + self.concept_read_encoder_shared_source_norm = nn.LayerNorm( + config.hidden_size, + eps=getattr(config, 'layernorm_epsilon', 1e-6), + dtype=dtype, + device=device) + for param in self.concept_read_encoder_shared_source_norm.parameters(): + param.requires_grad_(False) + + def set_attention_window(self, window_size, skip_frequency): + """Reference-compatible API. + + LMDeploy's rewrite bakes per-layer sliding-window policy into + ``ConceptLMV22VQOlmoBlock`` at construction time, so this is retained as + an explicit no-op for call-site compatibility. + """ + self._window_size = window_size + self._window_skip_frequency = skip_frequency + + def normalize_encoder_concept_states(self, encoder_concept_states: torch.Tensor): + """Apply the shared source norm used before concept-read-encoder routes.""" + return self.concept_read_encoder_shared_source_norm(encoder_concept_states) + + def predict_logits(self, hidden_states: torch.Tensor): + """Return logits in shape ``[..., num_codebooks, codebook_size]``.""" + return self.prediction_heads(hidden_states) + + def make_history_buffer(self, hidden_states: torch.Tensor) -> torch.Tensor: + """Allocate concept self-DD history buffer.""" + return self.concept_self_dd.make_history_buffer(hidden_states) + + @staticmethod + def write_history(history_buffer: torch.Tensor, slot_idx: int, hidden_states: torch.Tensor): + """Copy one concept history block into a layer-major history buffer.""" + return ConceptLMV22VQSelfDD.write_history(history_buffer, slot_idx, hidden_states) + + def forward(self, + concept_hidden: torch.Tensor, + encoder_concept_states: torch.Tensor, + position_ids: torch.Tensor, + past_key_values: list[list[torch.Tensor]] | None = None, + attn_metadata: Any = None, + encoder_source_dim: int = -2): + """Concept predictor forward. + + This mirrors the reference control flow, but uses LMDeploy's OLMo layer + rewrite and paged attention inputs. Full end-to-end use requires the + caller to provide concept-stream ``past_key_values`` and + ``attn_metadata`` matching the concept token layout. + """ + hidden_states = concept_hidden + history_buffer = self.make_history_buffer(hidden_states) + self.write_history(history_buffer, 0, hidden_states) + raw_states = [] + rotary_pos_emb = self.hlm_block._make_rotary_pos_emb(hidden_states, position_ids) + + for layer_idx, layer in enumerate(self.hlm_block.layers): + pkv = past_key_values[layer_idx] if past_key_values is not None else None + raw = layer( + hidden_states, + rotary_pos_emb=rotary_pos_emb, + past_key_value=pkv, + attn_metadata=attn_metadata, + ) + raw_states.append(raw) + self.write_history(history_buffer, layer_idx + 1, raw) + hidden_states = self.concept_self_dd.forward_from_buffer(layer_idx, raw, history_buffer) + hidden_states = self.concept_read_encoder_routes[layer_idx]( + hidden_states, + encoder_concept_states, + source_dim=encoder_source_dim, + ) + + hidden_states = self.hlm_block.final_layernorm(hidden_states) + logits = self.predict_logits(hidden_states) + return logits, raw_states + + +class ConceptLMV22VQOlmoAttention(nn.Module): + """Rewrite of ``_OlmoSelfAttention``. + + Differences from the reference: + - fused QKV uses lmdeploy's ``build_qkv_proj`` (standard [Q,K,V] packing, + TP-aware). The reference stores QKV per attention head, so the native + checkpoint weight is re-laid-out during ``load_weights``. + - q/k RMSNorm operates on the whole hidden_size (as in the reference), + not per-head. TP>1 needs an all-reduced variance; we follow internvl's + qkv_norm pattern. + - attention runs through lmdeploy's paged ``Attention`` primitive. + """ + + def __init__(self, + config: PretrainedConfig, + dtype: torch.dtype = None, + device: torch.device = None, + sliding_window: int | None = None, + prefix: str = ''): + super().__init__() + quantization_config = getattr(config, 'quantization_config', None) + hidden_size = config.hidden_size + num_heads = config.num_attention_heads + head_dim = config.kv_channels + num_kv_heads = num_heads + assert hidden_size == num_heads * head_dim, ( + f'ConceptLM OLMo attention expects hidden_size == num_heads * head_dim, ' + f'got {hidden_size} != {num_heads} * {head_dim}.') + + self.num_heads = num_heads + self.num_kv_heads = num_kv_heads + self.head_dim = head_dim + self.hidden_size = hidden_size + self.layernorm_epsilon = config.layernorm_epsilon + + # packed qkv + self.qkv_proj = build_qkv_proj( + hidden_size, + num_q_heads=num_heads, + num_kv_heads=num_kv_heads, + head_size=head_dim, + bias=False, + quant_config=quantization_config, + dtype=dtype, + device=device, + prefix=add_prefix('qkv_proj', prefix), + ) + + # q, k norm over the whole hidden_size (NOT per-head). tp=True with + # head_dim alignment so the weight shards correctly under TP. + self.q_layernorm = RMSNorm(hidden_size, + config.layernorm_epsilon, + quant_config=quantization_config, + dtype=dtype, + device=device, + tp=True, + align=head_dim, + prefix=add_prefix('q_layernorm', prefix)) + self.k_layernorm = RMSNorm(hidden_size, + config.layernorm_epsilon, + quant_config=quantization_config, + dtype=dtype, + device=device, + tp=True, + align=head_dim, + prefix=add_prefix('k_layernorm', prefix)) + + # rotary embedding + self.apply_rotary_pos_emb = ApplyRotaryEmb() + + # attention + self.attn_fwd = Attention(num_heads, + head_dim, + num_kv_heads=num_kv_heads, + v_head_size=head_dim, + sliding_window=None if sliding_window is None else int(sliding_window), + device=device) + + # o_proj + self.o_proj = build_o_proj(num_heads * head_dim, + hidden_size, + bias=False, + quant_config=quantization_config, + dtype=dtype, + device=device, + is_tp=True, + prefix=add_prefix('o_proj', prefix)) + + def _qkv_norm(self, q: torch.Tensor, k: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """RMSNorm over the whole hidden_size, TP-correct via all-reduce.""" + import lmdeploy.pytorch.distributed as dist + q_shape = q.shape + k_shape = k.shape + q = q.flatten(-2, -1) + k = k.flatten(-2, -1) + + tp, _ = dist.get_tp_world_rank('attn') + if tp == 1: + q = self.q_layernorm(q) + k = self.k_layernorm(k) + return q.view(q_shape), k.view(k_shape) + + # variance is computed over the full hidden_size, so it must be + # all-reduced across TP ranks before normalizing the local shard. + variance = _qk_rmsnorm_variance(q, k) + dist.all_reduce(variance) + q, k = _qk_rmsnorm_apply(q, k, variance, self.q_layernorm.weight, + self.k_layernorm.weight, self.hidden_size, + self.layernorm_epsilon) + return q.view(q_shape), k.view(k_shape) + + def forward(self, + hidden_states: torch.Tensor, + rotary_pos_emb: tuple[torch.Tensor, torch.Tensor], + past_key_value: list[torch.Tensor] | None = None, + attn_metadata: Any = None): + """Rewrite of _OlmoSelfAttention.forward.""" + # qkv proj -> (batch, seq, num_heads, head_dim) each + qkv_states = self.qkv_proj(hidden_states) + qkv_states = qkv_states.flatten(0, -2) # (-1, heads_total, head_dim) + query_states, key_states, value_states = self.qkv_proj.split_qkv(qkv_states) + + # q, k norm (whole hidden_size) + query_states, key_states = self._qkv_norm(query_states, key_states) + + # rotary embedding + cos, sin = rotary_pos_emb + query_states, key_states = self.apply_rotary_pos_emb( + query_states, key_states, cos, sin, inplace=True) + + # attention (paged) + attn_output = self.attn_fwd( + query_states, + key_states, + value_states, + past_key_value[0], + past_key_value[1], + attn_metadata, + k_scales_zeros=None if len(past_key_value) == 2 else past_key_value[2], + v_scales_zeros=None if len(past_key_value) == 2 else past_key_value[3], + inplace=True, + ) + attn_output = attn_output.reshape(*hidden_states.shape[:-1], -1) + + # o proj + attn_output = self.o_proj(attn_output) + return attn_output + + +class ConceptLMV22VQOlmoMLP(nn.Module): + """Rewrite of ``_OlmoMLP`` (SwiGLU).""" + + def __init__(self, + config: PretrainedConfig, + dtype: torch.dtype = None, + device: torch.device = None, + prefix: str = ''): + super().__init__() + quantization_config = getattr(config, 'quantization_config', None) + self.gate_up_proj = build_gateup_linear( + config.hidden_size, + [config.ffn_hidden_size, config.ffn_hidden_size], + bias=False, + dtype=dtype, + device=device, + quant_config=quantization_config, + is_tp=True, + prefix=add_prefix('gate_up_proj', prefix), + ) + self.act_fn = SiluAndMul(inplace=True) + self.down_proj = build_down_linear( + config.ffn_hidden_size, + config.hidden_size, + bias=False, + quant_config=quantization_config, + dtype=dtype, + device=device, + is_tp=True, + prefix=add_prefix('down_proj', prefix), + ) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + """forward.""" + gate_up = self.gate_up_proj(hidden_states) + act = self.act_fn(gate_up) + return self.down_proj(act) + + +class ConceptLMV22VQOlmoLayer(nn.Module): + """Rewrite of ``_OlmoLayer``. + + The reference uses post-norm residuals:: + h = h + post_attention_layernorm(attn(h)) + h = h + post_feedforward_layernorm(mlp(h)) + This block therefore calls RMSNorm without the residual argument and adds + the residual explicitly. + """ + + def __init__(self, + config: PretrainedConfig, + layer_idx: int, + dtype: torch.dtype = None, + device: torch.device = None, + sliding_window: int | None = None, + prefix: str = ''): + super().__init__() + quantization_config = getattr(config, 'quantization_config', None) + self.layer_number = int(layer_idx) + 1 + self.self_attention = ConceptLMV22VQOlmoAttention(config, + dtype=dtype, + device=device, + sliding_window=sliding_window, + prefix=add_prefix('self_attention', prefix)) + self.post_attention_layernorm = RMSNorm(config.hidden_size, + config.layernorm_epsilon, + quant_config=quantization_config, + dtype=dtype, + device=device, + prefix=add_prefix('post_attention_layernorm', prefix)) + self.mlp = ConceptLMV22VQOlmoMLP(config, dtype=dtype, device=device, prefix=add_prefix('mlp', prefix)) + self.post_feedforward_layernorm = RMSNorm(config.hidden_size, + config.layernorm_epsilon, + quant_config=quantization_config, + dtype=dtype, + device=device, + prefix=add_prefix('post_feedforward_layernorm', prefix)) + + def forward(self, + hidden_states: torch.Tensor, + rotary_pos_emb: tuple[torch.Tensor, torch.Tensor], + past_key_value: list[torch.Tensor] | None = None, + attn_metadata: Any = None): + """forward. + + The reference uses post-norm residuals:: + h = h + post_attention_layernorm(attn(h)) + h = h + post_feedforward_layernorm(mlp(h)) + This is NOT the same as lmdeploy's pre-norm residual form + ``norm(x, residual=h)`` (which normalizes x+h). So we call the norm + without a residual and add explicitly. + """ + attn_out = self.self_attention( + hidden_states=hidden_states, + rotary_pos_emb=rotary_pos_emb, + past_key_value=past_key_value, + attn_metadata=attn_metadata, + ) + hidden_states = hidden_states + self.post_attention_layernorm(attn_out) + + mlp_out = self.mlp(hidden_states) + hidden_states = hidden_states + self.post_feedforward_layernorm(mlp_out) + return hidden_states + + +class ConceptLMV22VQOlmoBlock(nn.Module): + """Rewrite of ``_OlmoBlock``. + + Holds a stack of ``_OlmoLayer`` and an optional final RMSNorm. Mirrors the + reference's ``forward`` return contract: ``(hidden_states, layer_states)``. + """ + + def __init__(self, + config: PretrainedConfig, + num_layers: int, + post_layer_norm: bool, + dtype: torch.dtype = None, + device: torch.device = None, + window_size: Any = _CONFIG_VALUE, + skip_frequency: Any = _CONFIG_VALUE, + prefix: str = ''): + super().__init__() + quantization_config = getattr(config, 'quantization_config', None) + if window_size is _CONFIG_VALUE: + window_size = _get_configured_window(config) + elif isinstance(window_size, (list, tuple)): + window_size = window_size[0] if len(window_size) > 0 else None + if window_size is not None: + window_size = int(window_size) + if skip_frequency is _CONFIG_VALUE: + skip_frequency = getattr(config, 'window_attn_skip_freq', None) + if skip_frequency is not None: + skip_frequency = int(skip_frequency) + + self.layers = nn.ModuleList([ + ConceptLMV22VQOlmoLayer(config, + layer_idx, + dtype=dtype, + device=device, + sliding_window=self._layer_sliding_window(layer_idx + 1, window_size, + skip_frequency), + prefix=add_prefix(f'layers.{layer_idx}', prefix)) + for layer_idx in range(num_layers) + ]) + self.final_layernorm = ( + RMSNorm(config.hidden_size, + config.layernorm_epsilon, + quant_config=quantization_config, + dtype=dtype, + device=device, + prefix=add_prefix('final_layernorm', prefix)) if post_layer_norm else None) + self.rotary_emb = _make_olmo_rotary_embedding(config, device=device) + self.num_heads = int(config.num_attention_heads) + self.head_dim = int(config.kv_channels) + + @staticmethod + def _layer_sliding_window(layer_number: int, window_size: int | None, skip_frequency: int | None): + """Match reference OLMo window/full attention alternation.""" + if window_size is None: + return None + if skip_frequency is not None and layer_number % skip_frequency == 0: + return None + return window_size + + def _make_rotary_pos_emb(self, hidden_states: torch.Tensor, position_ids: torch.Tensor): + """Create RoPE cos/sin in the shape consumed by ``ApplyRotaryEmb``.""" + if position_ids.dim() == 1: + position_ids = position_ids.unsqueeze(0) + cos, sin = self.rotary_emb(hidden_states, position_ids) + if cos.dim() == 3 and cos.size(0) == 1: + cos = cos[0] + sin = sin[0] + return cos, sin + + def forward(self, + hidden_states: torch.Tensor, + position_ids: torch.Tensor, + past_key_values: list[list[torch.Tensor]] | None = None, + attn_metadata: Any = None, + collect: bool = False): + """forward.""" + layer_states = [] + rotary_pos_emb = self._make_rotary_pos_emb(hidden_states, position_ids) + for idx, layer in enumerate(self.layers): + pkv = past_key_values[idx] if past_key_values is not None else None + hidden_states = layer( + hidden_states, + rotary_pos_emb=rotary_pos_emb, + past_key_value=pkv, + attn_metadata=attn_metadata, + ) + if collect: + layer_states.append(hidden_states) + if self.final_layernorm is not None: + hidden_states = self.final_layernorm(hidden_states) + return hidden_states, layer_states + + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]], prefix: str = ''): + """Load native ConceptLM OLMo block weights into the LMDeploy rewrite.""" + if prefix and not prefix.endswith('.'): + prefix = f'{prefix}.' + + params_dict = dict(self.named_parameters()) + loaded_names = set() + for name, loaded_weight in weights: + if 'rotary_emb.inv_freq' in name: + continue + if prefix: + if not name.startswith(prefix): + continue + name = name[len(prefix):] + + if name.endswith('.self_attention.qkv_proj.weight'): + param = params_dict[name] + query, key, value = param.weight_spliter(loaded_weight) + load_weight(param, query, shard_id='q') + load_weight(param, key, shard_id='k') + load_weight(param, value, shard_id='v') + loaded_names.add(name) + continue + + if name.endswith('.self_attention.linear_qkv.weight'): + target_name = name.replace('.self_attention.linear_qkv.weight', + '.self_attention.qkv_proj.weight') + param = params_dict[target_name] + loaded_weight = _repack_olmo_qkv_weight(loaded_weight, self.num_heads, self.head_dim) + query, key, value = param.weight_spliter(loaded_weight) + load_weight(param, query, shard_id='q') + load_weight(param, key, shard_id='k') + load_weight(param, value, shard_id='v') + loaded_names.add(target_name) + continue + + if name.endswith('.self_attention.linear_proj.weight'): + target_name = name.replace('.self_attention.linear_proj.weight', + '.self_attention.o_proj.weight') + elif name.endswith('.mlp.gate_up_proj.weight'): + param = params_dict[name] + gate, up = param.weight_spliter(loaded_weight) + load_weight(param, gate, shard_id=0) + load_weight(param, up, shard_id=1) + loaded_names.add(name) + continue + elif name.endswith('.mlp.linear_fc1.weight'): + target_name = name.replace('.mlp.linear_fc1.weight', '.mlp.gate_up_proj.weight') + param = params_dict[target_name] + gate, up = param.weight_spliter(loaded_weight) + load_weight(param, gate, shard_id=0) + load_weight(param, up, shard_id=1) + loaded_names.add(target_name) + continue + elif name.endswith('.mlp.linear_fc2.weight'): + target_name = name.replace('.mlp.linear_fc2.weight', '.mlp.down_proj.weight') + else: + target_name = name + + param = params_dict.get(target_name) + if param is None: + continue + load_weight(param, loaded_weight) + loaded_names.add(target_name) + return loaded_names + + +class ConceptLMV22VQForCausalLM(nn.Module, DeployModelMixinV1, CudaGraphMixin): + """Rewrote model of ConceptLMV22VQForCausalLM.""" + + def __init__(self, + config: PretrainedConfig, + ctx_mgr: StepContextManager, + dtype: torch.dtype = None, + device: torch.device = None, + prefix: str = ''): + super().__init__() + self.config = config + self.ctx_mgr = ctx_mgr + # token embedding — mirrors ``self.embedding`` in the reference. + self.embedding = ConceptLMV22VQEmbedding(config, dtype=dtype, device=device) + self.concept_quantizer = ConceptLMV22VQQuantizer(config, dtype=dtype, device=device) + self.dd_encoder_self_dd = ConceptLMV22VQSelfDD(config, + config.concept_encoder_layers, + use_softmax=False, + dtype=dtype, + device=device) + self.decoder_read_encoder_routes = nn.ModuleList([ + ConceptLMV22VQResidualRoute(config, + config.concept_encoder_layers, + use_softmax=True, + dtype=dtype, + device=device) + for _ in range(config.concept_decoder_layers) + ]) + self.decoder_read_concept_routes = nn.ModuleList([ + ConceptLMV22VQResidualRoute(config, + config.concept_special_layers, + use_softmax=True, + dtype=dtype, + device=device) + for _ in range(config.concept_decoder_layers) + ]) + self.dd_two_route_add = ConceptLMV22VQTwoRouteAdd(config, dtype=dtype, device=device) + self.concept_predictor = ConceptLMV22VQConceptPredictor(config, + dtype=dtype, + device=device, + prefix=add_prefix('concept_predictor', prefix)) + self.concept_predictor.set_attention_window(tuple(getattr(config, 'window_size', (None, None))), + getattr(config, 'window_attn_skip_freq', None)) + # output projection — mirrors ``output_layer`` in the reference. Built + # via build_lm_head and named ``lm_head`` so DeployModelMixinV1. + # get_logits picks it up directly; load_weights maps the checkpoint's + # ``output_layer`` onto it. + self.lm_head = self.build_lm_head( + config.hidden_size, config.vocab_size, bias=False, dtype=dtype, device=device) + + def forward(self, + input_ids: torch.Tensor, + position_ids: torch.Tensor, + past_key_values: list[list[torch.Tensor]], + attn_metadata=None, + inputs_embeds: torch.Tensor = None, + **kwargs): + """Model forward, return hidden_states (logits computed by runtime).""" + if inputs_embeds is None: + # NOTE: placeholder. The real path is + # embed -> encoder -> concept(vq+predictor) -> fusion -> decoder + # added incrementally. Returns raw embeddings as hidden_states. + hidden_states = self.embedding(input_ids) + else: + hidden_states = inputs_embeds + return hidden_states + + def get_input_embeddings(self): + """Get input embeddings.""" + return self.embedding.word_embeddings + + def prepare_inputs_for_generation(self, + past_key_values: list[list[torch.Tensor]], + inputs_embeds: torch.Tensor | None = None, + context: StepContext = None): + """Prepare input.""" + input_ids = context.input_ids + position_ids = context.position_ids + attn_metadata = context.attn_metadata + return dict( + input_ids=input_ids, + position_ids=position_ids, + past_key_values=past_key_values, + attn_metadata=attn_metadata, + inputs_embeds=inputs_embeds, + ) + + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]): + """Load weights. + + Only modules currently wired into the top-level model are loaded. Other + checkpoint tensors are skipped until their corresponding modules are + added. + """ + # (checkpoint_name, target_name) + weight_map = { + 'embedding.word_embeddings.weight': 'embedding.word_embeddings.weight', + 'output_layer.weight': 'lm_head.weight', + } + codebook_prefix = 'concept_quantizer.codebook.' + concept_hlm_prefix = 'concept_predictor.hlm_block.' + prediction_head_prefix = 'concept_predictor.prediction_heads.' + params_dict = dict(self.named_parameters()) + for name, loaded_weight in weights: + if 'rotary_emb.inv_freq' in name: + continue + if name.startswith(concept_hlm_prefix): + self.concept_predictor.hlm_block.load_weights( + [(name, loaded_weight)], + prefix=concept_hlm_prefix[:-1], + ) + continue + if name.startswith(prediction_head_prefix): + suffix = name[len(prediction_head_prefix):] + parts = suffix.split('.') + if len(parts) == 2 and parts[0].isdigit() and parts[1] in ('weight', 'bias'): + target_name = f'{prediction_head_prefix}proj.{parts[1]}' + param = params_dict.get(target_name) + if param is not None: + load_weight(param, loaded_weight, shard_id=int(parts[0])) + continue + if suffix in ('proj.weight', 'proj.bias'): + param = params_dict.get(name) + if param is not None: + for shard_id, shard_weight in enumerate(param.weight_spliter(loaded_weight)): + load_weight(param, shard_weight, shard_id=shard_id) + continue + if name.startswith(codebook_prefix): + codebook_idx = int(name[len(codebook_prefix):]) + param = params_dict['concept_quantizer.codebook'] + _load_stacked_codebook_weight(param, loaded_weight, codebook_idx) + continue + target = weight_map.get(name) + if target is None: + target = name + if target not in params_dict: + # skip modules not yet wired into the top-level model + continue + param = params_dict[target] + load_weight(param, loaded_weight) diff --git a/lmdeploy/pytorch/models/module_map.py b/lmdeploy/pytorch/models/module_map.py index 27e74e8e67..8e0cc839d8 100644 --- a/lmdeploy/pytorch/models/module_map.py +++ b/lmdeploy/pytorch/models/module_map.py @@ -265,6 +265,11 @@ 'InternLM3ForCausalLM': f'{LMDEPLOY_PYTORCH_MODEL_PATH}.internlm3.InternLM3ForCausalLM', }) +# conceptlm v22 vq (intern ncp) +MODULE_MAP.update({ + 'ConceptLMV22VQForCausalLM': f'{LMDEPLOY_PYTORCH_MODEL_PATH}.intern_ncp.ConceptLMV22VQForCausalLM', +}) + # internlm2 reward model MODULE_MAP.update( {'InternLM2ForRewardModel': f'{LMDEPLOY_PYTORCH_MODEL_PATH}.internlm2_reward.InternLM2ForRewardModel'}) From 2da45c29354bb8d9e982e01fa9acc56230508019 Mon Sep 17 00:00:00 2001 From: grimoire Date: Fri, 24 Jul 2026 13:13:00 +0800 Subject: [PATCH 02/16] Add ConceptLM runtime cache ops --- lmdeploy/pytorch/backends/base.py | 3 + lmdeploy/pytorch/backends/conceptlm.py | 36 ++ lmdeploy/pytorch/backends/cuda/conceptlm.py | 38 ++ lmdeploy/pytorch/backends/cuda/op_backend.py | 3 + .../pytorch/backends/default/conceptlm.py | 88 +++ .../pytorch/backends/default/op_backend.py | 3 + lmdeploy/pytorch/configurations/conceptlm.py | 41 ++ lmdeploy/pytorch/kernels/cuda/conceptlm.py | 167 ++++++ lmdeploy/pytorch/models/intern_ncp.py | 539 +++++++++++++++++- lmdeploy/pytorch/nn/__init__.py | 1 + lmdeploy/pytorch/nn/conceptlm.py | 56 ++ 11 files changed, 953 insertions(+), 22 deletions(-) create mode 100644 lmdeploy/pytorch/backends/conceptlm.py create mode 100644 lmdeploy/pytorch/backends/cuda/conceptlm.py create mode 100644 lmdeploy/pytorch/backends/default/conceptlm.py create mode 100644 lmdeploy/pytorch/kernels/cuda/conceptlm.py create mode 100644 lmdeploy/pytorch/nn/conceptlm.py diff --git a/lmdeploy/pytorch/backends/base.py b/lmdeploy/pytorch/backends/base.py index 3ad683b980..04b4b1704d 100644 --- a/lmdeploy/pytorch/backends/base.py +++ b/lmdeploy/pytorch/backends/base.py @@ -45,6 +45,9 @@ class OpType(Enum): CausalConv1d = auto() GatedDeltaRule = auto() + # ConceptLM + ConceptLMRuntimeOps = auto() + class OpsBackend(ABC): """Layer backend abstract.""" diff --git a/lmdeploy/pytorch/backends/conceptlm.py b/lmdeploy/pytorch/backends/conceptlm.py new file mode 100644 index 0000000000..a02974c707 --- /dev/null +++ b/lmdeploy/pytorch/backends/conceptlm.py @@ -0,0 +1,36 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from abc import ABC, abstractmethod + +from torch import Tensor + + +class ConceptLMRuntimeOpsImpl(ABC): + """ConceptLM runtime operation implementation. + + Model-specific runtime/cache operations live behind this single backend + interface. That keeps model code free from direct kernel calls while + avoiding one OpType/nn module per small ConceptLM state operation. + """ + + @abstractmethod + def decode_chunk_state_update( + self, + chunk_source_state_cache: Tensor, + current_source_states: Tensor, + state_ids: Tensor, + position_ids: Tensor, + chunk_size: int, + merge_method: str, + ) -> tuple[Tensor, Tensor, Tensor]: + """Update state cache and return concept inputs, next rows, and mask.""" + raise NotImplementedError('Not implemented.') + + +class ConceptLMRuntimeOpsBuilder(ABC): + """ConceptLM runtime operation builder.""" + + @staticmethod + @abstractmethod + def build() -> ConceptLMRuntimeOpsImpl: + """Build layer implementation.""" + raise NotImplementedError('Not implemented.') diff --git a/lmdeploy/pytorch/backends/cuda/conceptlm.py b/lmdeploy/pytorch/backends/cuda/conceptlm.py new file mode 100644 index 0000000000..951e34ac96 --- /dev/null +++ b/lmdeploy/pytorch/backends/cuda/conceptlm.py @@ -0,0 +1,38 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from torch import Tensor + +from lmdeploy.pytorch.kernels.cuda.conceptlm import decode_chunk_state_update + +from ..conceptlm import ConceptLMRuntimeOpsBuilder, ConceptLMRuntimeOpsImpl + + +class TritonConceptLMRuntimeOpsImpl(ConceptLMRuntimeOpsImpl): + """Triton implementation of ConceptLM runtime operations.""" + + def decode_chunk_state_update( + self, + chunk_source_state_cache: Tensor, + current_source_states: Tensor, + state_ids: Tensor, + position_ids: Tensor, + chunk_size: int, + merge_method: str, + ) -> tuple[Tensor, Tensor, Tensor]: + """Update state cache and return fixed-shape decode rows.""" + return decode_chunk_state_update( + chunk_source_state_cache, + current_source_states, + state_ids, + position_ids, + chunk_size, + merge_method, + ) + + +class TritonConceptLMRuntimeOpsBuilder(ConceptLMRuntimeOpsBuilder): + """Triton ConceptLM runtime operation builder.""" + + @staticmethod + def build() -> ConceptLMRuntimeOpsImpl: + """Build layer implementation.""" + return TritonConceptLMRuntimeOpsImpl() diff --git a/lmdeploy/pytorch/backends/cuda/op_backend.py b/lmdeploy/pytorch/backends/cuda/op_backend.py index 67c29ae26d..1464472cba 100644 --- a/lmdeploy/pytorch/backends/cuda/op_backend.py +++ b/lmdeploy/pytorch/backends/cuda/op_backend.py @@ -91,6 +91,9 @@ def get_layer_impl_builder(cls, layer_type: OpType): elif layer_type == OpType.GatedDeltaRule: from .gated_delta_rule import CudaGatedDeltaRuleBuilder return CudaGatedDeltaRuleBuilder + elif layer_type == OpType.ConceptLMRuntimeOps: + from .conceptlm import TritonConceptLMRuntimeOpsBuilder + return TritonConceptLMRuntimeOpsBuilder else: logger.debug(f'Op {layer_type} fallback to default implementation.') return super().get_layer_impl_builder(layer_type) diff --git a/lmdeploy/pytorch/backends/default/conceptlm.py b/lmdeploy/pytorch/backends/default/conceptlm.py new file mode 100644 index 0000000000..558dc4f0f0 --- /dev/null +++ b/lmdeploy/pytorch/backends/default/conceptlm.py @@ -0,0 +1,88 @@ +# Copyright (c) OpenMMLab. All rights reserved. +import torch +from torch import Tensor + +from ..conceptlm import ConceptLMRuntimeOpsBuilder, ConceptLMRuntimeOpsImpl + + +def _flatten_decode_position_ids(position_ids: Tensor, batch_size: int, device: torch.device) -> Tensor: + """Normalize decode position ids to one absolute position per batch row.""" + if position_ids.dim() == 0: + position_ids = position_ids.view(1) + if position_ids.dim() == 1: + return position_ids.to(device=device, dtype=torch.long) + position_ids = position_ids.reshape(-1) + if position_ids.numel() == batch_size: + return position_ids.to(device=device, dtype=torch.long) + assert position_ids.numel() % batch_size == 0, ( + f'Cannot map position_ids with {position_ids.numel()} elements to batch size {batch_size}.') + return position_ids.reshape(-1, batch_size)[-1].to(device=device, dtype=torch.long) + + +class DefaultConceptLMRuntimeOpsImpl(ConceptLMRuntimeOpsImpl): + """Torch fallback implementation of ConceptLM runtime operations.""" + + def decode_chunk_state_update( + self, + chunk_source_state_cache: Tensor, + current_source_states: Tensor, + state_ids: Tensor, + position_ids: Tensor, + chunk_size: int, + merge_method: str, + ) -> tuple[Tensor, Tensor, Tensor]: + """Update state cache and return fixed-shape decode rows.""" + assert current_source_states.dim() == 3, ( + f'current_source_states must be [batch, num_sources, hidden], got {tuple(current_source_states.shape)}.') + assert chunk_source_state_cache.dim() == 3, ( + f'chunk_source_state_cache must be [num_state_slots, num_sources, hidden], ' + f'got {tuple(chunk_source_state_cache.shape)}.') + batch_size = current_source_states.size(0) + assert current_source_states.shape[1:] == chunk_source_state_cache.shape[1:], ( + f'Current source state shape {tuple(current_source_states.shape[1:])} does not match state-cache shape ' + f'{tuple(chunk_source_state_cache.shape[1:])}.') + + state_ids = state_ids.to(device=current_source_states.device, dtype=torch.long) + position_ids = _flatten_decode_position_ids(position_ids, batch_size, current_source_states.device) + assert position_ids.numel() == batch_size, ( + f'Expected {batch_size} decode position ids, got {position_ids.numel()}.') + + valid_state_mask = state_ids >= 0 + safe_state_ids = state_ids.clamp(min=0) + previous_rows = chunk_source_state_cache.index_select(0, safe_state_ids) + + chunk_size = int(chunk_size) + chunk_pos = torch.remainder(position_ids, chunk_size) + update_mask = valid_state_mask & (torch.remainder(position_ids + 1, chunk_size) == 0) + first_token_mask = valid_state_mask & (chunk_pos == 0) + merge_method = str(merge_method) + + if merge_method == 'first': + update_rows = torch.where(first_token_mask.view(batch_size, 1, 1), current_source_states, previous_rows) + concept_input_states = update_rows + elif merge_method == 'last': + update_rows = current_source_states + concept_input_states = current_source_states + else: + update_rows = previous_rows + current_source_states + concept_input_states = update_rows / chunk_size + + zero_rows = torch.zeros_like(update_rows) + next_rows = torch.where(update_mask.view(batch_size, 1, 1), zero_rows, update_rows) + next_rows = torch.where(valid_state_mask.view(batch_size, 1, 1), next_rows, previous_rows) + concept_input_states = torch.where(update_mask.view(batch_size, 1, 1), concept_input_states, zero_rows) + + for batch_idx in range(batch_size): + state_id = int(state_ids[batch_idx]) + if state_id >= 0: + chunk_source_state_cache[state_id].copy_(next_rows[batch_idx]) + return concept_input_states, next_rows, update_mask + + +class DefaultConceptLMRuntimeOpsBuilder(ConceptLMRuntimeOpsBuilder): + """Torch fallback ConceptLM runtime operation builder.""" + + @staticmethod + def build() -> ConceptLMRuntimeOpsImpl: + """Build layer implementation.""" + return DefaultConceptLMRuntimeOpsImpl() diff --git a/lmdeploy/pytorch/backends/default/op_backend.py b/lmdeploy/pytorch/backends/default/op_backend.py index 6bfb9e5934..22a530cbfc 100644 --- a/lmdeploy/pytorch/backends/default/op_backend.py +++ b/lmdeploy/pytorch/backends/default/op_backend.py @@ -50,6 +50,9 @@ def get_layer_impl_builder(cls, layer_type: OpType): elif layer_type == OpType.RouterNoauxTC: from .moe_router import DefaultRouterNoauxTCBuilder return DefaultRouterNoauxTCBuilder + elif layer_type == OpType.ConceptLMRuntimeOps: + from .conceptlm import DefaultConceptLMRuntimeOpsBuilder + return DefaultConceptLMRuntimeOpsBuilder else: raise RuntimeError(f'{layer_type} not supported.') diff --git a/lmdeploy/pytorch/configurations/conceptlm.py b/lmdeploy/pytorch/configurations/conceptlm.py index b16210249b..c4e5273728 100644 --- a/lmdeploy/pytorch/configurations/conceptlm.py +++ b/lmdeploy/pytorch/configurations/conceptlm.py @@ -1,4 +1,6 @@ # Copyright (c) OpenMMLab. All rights reserved. +import torch + from lmdeploy.utils import get_logger from .builder import AutoModelConfigBuilder @@ -6,6 +8,28 @@ logger = get_logger('lmdeploy') +CONCEPT_STATE_CHUNK_SOURCE = 0 +CONCEPT_STATE_LAST_RAW = 1 +CONCEPT_STATE_LAST_FINAL = 2 +CONCEPT_STATE_NAMES = ( + 'concept_chunk_source_state', + 'concept_last_raw_states', + 'concept_last_final_state', +) + + +def _get_concept_state_dtype(hf_config): + """Return the dtype used by ConceptLM sequence-state caches.""" + torch_dtype = getattr(hf_config, 'torch_dtype', None) + if isinstance(torch_dtype, torch.dtype): + return torch_dtype + torch_dtype = str(torch_dtype).lower() + if 'bfloat16' in torch_dtype or 'bf16' in torch_dtype: + return torch.bfloat16 + if 'float32' in torch_dtype or 'fp32' in torch_dtype: + return torch.float32 + return torch.float16 + class ConceptLMModelConfigBuilder(AutoModelConfigBuilder): """Config builder for ConceptLM V2.2-VQ. @@ -44,6 +68,23 @@ def build(cls, hf_config, model_path: str = None, **kwargs): model_config.llm_config.concept_kv_concept_offset = enc_layers model_config.llm_config.concept_kv_decoder_offset = enc_layers + concept_layers model_config.llm_config.concept_kv_total_layers = model_config.num_layers + + hidden_size = int(hf_config.hidden_size) + state_dtype = _get_concept_state_dtype(hf_config) + concept_encoder_read_sources = max(enc_layers - 1, 0) + model_config.states_shapes = [ + ((concept_encoder_read_sources, hidden_size), state_dtype), + ((concept_layers, hidden_size), state_dtype), + ((hidden_size, ), state_dtype), + ] + # The current branch only supports anonymous states_shapes. Keep stable + # indices on the HF config so model code has one semantic source of + # truth; if DSV4 StateCacheSpec lands here later, these become the + # names of ConceptLM's state-cache specs. + model_config.llm_config.concept_state_names = CONCEPT_STATE_NAMES + model_config.llm_config.concept_state_chunk_source_idx = CONCEPT_STATE_CHUNK_SOURCE + model_config.llm_config.concept_state_last_raw_idx = CONCEPT_STATE_LAST_RAW + model_config.llm_config.concept_state_last_final_idx = CONCEPT_STATE_LAST_FINAL return model_config @staticmethod diff --git a/lmdeploy/pytorch/kernels/cuda/conceptlm.py b/lmdeploy/pytorch/kernels/cuda/conceptlm.py new file mode 100644 index 0000000000..aaee94c7db --- /dev/null +++ b/lmdeploy/pytorch/kernels/cuda/conceptlm.py @@ -0,0 +1,167 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""ConceptLM runtime kernels.""" + +import torch +import triton +import triton.language as tl + + +@triton.jit +def _decode_chunk_state_update_kernel( + state_cache, + current_states, + state_ids, + position_ids, + concept_inputs, + next_rows, + update_mask, + state_stride_n, + state_stride_s, + state_stride_h, + cur_stride_b, + cur_stride_s, + cur_stride_h, + out_stride_b, + out_stride_s, + out_stride_h, + next_stride_b, + next_stride_s, + next_stride_h, + HIDDEN: tl.constexpr, + TOTAL_ELEMS: tl.constexpr, + CHUNK_SIZE: tl.constexpr, + MERGE_METHOD: tl.constexpr, + BLOCK: tl.constexpr, +): + tile_id = tl.program_id(0) + batch_id = tl.program_id(1) + offs = tile_id * BLOCK + tl.arange(0, BLOCK) + valid_elem = offs < TOTAL_ELEMS + + source_id = offs // HIDDEN + hidden_id = offs - source_id * HIDDEN + + state_id = tl.load(state_ids + batch_id) + valid_state = state_id >= 0 + safe_state_id = tl.maximum(state_id, 0) + pos = tl.load(position_ids + batch_id) + chunk_pos = pos % CHUNK_SIZE + is_boundary = ((pos + 1) % CHUNK_SIZE) == 0 + is_first_token = chunk_pos == 0 + + current_ptrs = (current_states + batch_id * cur_stride_b + source_id * cur_stride_s + + hidden_id * cur_stride_h) + state_ptrs = state_cache + safe_state_id * state_stride_n + source_id * state_stride_s + hidden_id * state_stride_h + + current = tl.load(current_ptrs, mask=valid_elem, other=0.0).to(tl.float32) + previous = tl.load(state_ptrs, mask=valid_elem, other=0.0).to(tl.float32) + + if MERGE_METHOD == 1: # first + merged = tl.where(is_first_token, current, previous) + concept = merged + elif MERGE_METHOD == 2: # last + merged = current + concept = current + else: # meanpooling + merged = previous + current + concept = merged / CHUNK_SIZE + + zero = tl.zeros((BLOCK, ), dtype=tl.float32) + next_value = tl.where(is_boundary, zero, merged) + concept_value = tl.where(valid_state & is_boundary, concept, zero) + next_debug_value = tl.where(valid_state, next_value, previous) + + concept_ptrs = concept_inputs + batch_id * out_stride_b + source_id * out_stride_s + hidden_id * out_stride_h + next_ptrs = next_rows + batch_id * next_stride_b + source_id * next_stride_s + hidden_id * next_stride_h + tl.store(concept_ptrs, concept_value, mask=valid_elem) + tl.store(next_ptrs, next_debug_value, mask=valid_elem) + tl.store(state_ptrs, next_value, mask=valid_elem & valid_state) + + if tile_id == 0: + tl.store(update_mask + batch_id, valid_state & is_boundary) + + +def _flatten_decode_position_ids(position_ids: torch.Tensor, batch_size: int, device: torch.device) -> torch.Tensor: + """Normalize decode position ids to one absolute position per batch row.""" + if position_ids.dim() == 0: + position_ids = position_ids.view(1) + if position_ids.dim() == 1: + return position_ids.to(device=device, dtype=torch.long) + position_ids = position_ids.reshape(-1) + if position_ids.numel() == batch_size: + return position_ids.to(device=device, dtype=torch.long) + assert position_ids.numel() % batch_size == 0, ( + f'Cannot map position_ids with {position_ids.numel()} elements to batch size {batch_size}.') + return position_ids.reshape(-1, batch_size)[-1].to(device=device, dtype=torch.long) + + +def _merge_method_id(merge_method: str) -> int: + """Map ConceptLM merge method string to kernel constexpr id.""" + merge_method = str(merge_method) + if merge_method == 'first': + return 1 + if merge_method == 'last': + return 2 + return 0 + + +def decode_chunk_state_update( + chunk_source_state_cache: torch.Tensor, + current_source_states: torch.Tensor, + state_ids: torch.Tensor, + position_ids: torch.Tensor, + chunk_size: int, + merge_method: str, + block: int = 1024, +): + """Update ConceptLM decode chunk-source state in-place. + + Args: + chunk_source_state_cache: ``[num_state_slots, num_sources, hidden]``. + current_source_states: ``[batch, num_sources, hidden]``. + state_ids: ``[batch]`` with ``-1`` for padded graph rows. + position_ids: absolute decode positions. + chunk_size: ConceptLM chunk size. + merge_method: ``meanpooling``, ``first``, or ``last``. + block: Triton vector width. + + Returns: + Tuple ``(concept_inputs, next_rows, update_mask)``. ``concept_inputs`` + is zero for non-boundary rows. ``next_rows`` is a debug/reference copy + of the per-batch rows written to state cache. ``update_mask`` is + ``True`` only for valid boundary rows. + """ + assert chunk_source_state_cache.is_cuda, 'ConceptLM chunk-state kernel requires CUDA state cache.' + assert current_source_states.is_cuda, 'ConceptLM chunk-state kernel requires CUDA current states.' + assert current_source_states.dim() == 3 + assert chunk_source_state_cache.dim() == 3 + assert current_source_states.shape[1:] == chunk_source_state_cache.shape[1:] + + batch_size, num_sources, hidden = current_source_states.shape + total_elems = num_sources * hidden + state_ids = state_ids.to(device=current_source_states.device, dtype=torch.long) + position_ids = _flatten_decode_position_ids(position_ids, batch_size, current_source_states.device) + concept_inputs = torch.empty_like(current_source_states) + next_rows = torch.empty_like(current_source_states) + update_mask = torch.empty((batch_size, ), dtype=torch.bool, device=current_source_states.device) + grid = (triton.cdiv(total_elems, block), batch_size) + _decode_chunk_state_update_kernel[grid]( + chunk_source_state_cache, + current_source_states, + state_ids, + position_ids, + concept_inputs, + next_rows, + update_mask, + *chunk_source_state_cache.stride(), + *current_source_states.stride(), + *concept_inputs.stride(), + *next_rows.stride(), + HIDDEN=hidden, + TOTAL_ELEMS=total_elems, + CHUNK_SIZE=int(chunk_size), + MERGE_METHOD=_merge_method_id(merge_method), + BLOCK=block, + num_warps=8, + ) + return concept_inputs, next_rows, update_mask diff --git a/lmdeploy/pytorch/models/intern_ncp.py b/lmdeploy/pytorch/models/intern_ncp.py index 9feaa17ad1..41c439e3c4 100644 --- a/lmdeploy/pytorch/models/intern_ncp.py +++ b/lmdeploy/pytorch/models/intern_ncp.py @@ -14,13 +14,16 @@ - ``_ResidualRoute``: replicated residual source mixer for decoder routes. - ``_TwoRouteAdd``: decoder depth mixing plus final-concept residual route. - ``_ConceptPredictor``: concept block container and prediction heads. + - top-level encoder/decoder containers, fusion norms, route norms, and + checkpoint loading for the implemented module tree. Results are NOT yet correct end-to-end — the encoder/concept/decoder control -flow is added in later steps. The module structure mirrors the reference for -readability. +flow still needs runtime metadata for the compressed concept stream. The +module structure mirrors the reference for readability. """ from collections.abc import Iterable +from dataclasses import dataclass from typing import Any import torch @@ -29,7 +32,8 @@ from transformers.configuration_utils import PretrainedConfig from lmdeploy.pytorch.model_inputs import StepContext, StepContextManager -from lmdeploy.pytorch.nn import ApplyRotaryEmb, Attention, RMSNorm, SiluAndMul, build_rotary_embedding +from lmdeploy.pytorch.nn import (ApplyRotaryEmb, Attention, ConceptLMRuntimeOps, RMSNorm, SiluAndMul, + build_rotary_embedding) from lmdeploy.pytorch.nn.linear import (build_down_linear, build_gateup_linear, build_merged_colwise_linear, build_o_proj, build_qkv_proj) from lmdeploy.pytorch.weight_loader.model_weight_loader import load_weight @@ -43,6 +47,115 @@ _SourceStates = list[torch.Tensor] | tuple[torch.Tensor, ...] | torch.Tensor | None +@dataclass +class ConceptMetadata: + """Layer-invariant ConceptLM runtime metadata. + + This mirrors the DSV4 pattern: ``StepContext`` is read at the top-level + model boundary, then submodules receive explicit metadata instead of + reaching back into the engine context. The dense/reference helpers below do + not consume all fields yet; they are part of the serving decode contract. + """ + + chunk_size: int + merge_method: str + shift_feature: bool + is_decoding: bool | None = None + state_ids: torch.Tensor | None = None + position_ids: torch.Tensor | None = None + block_offsets: torch.Tensor | None = None + q_seqlens: torch.Tensor | None = None + kv_seqlens: torch.Tensor | None = None + q_start_loc: torch.Tensor | None = None + attn_metadata: Any = None + + @classmethod + def build(cls, + config: PretrainedConfig, + position_ids: torch.Tensor | None = None, + attn_metadata: Any = None, + state_ids: torch.Tensor | None = None): + """Build ConceptLM metadata from explicit forward inputs.""" + return cls( + chunk_size=int(config.concept_chunk_size), + merge_method=getattr(config, 'concept_chunk_merge_method', 'meanpooling'), + shift_feature=bool(getattr(config, 'concept_shift_feature', True)), + is_decoding=getattr(attn_metadata, 'is_decoding', None), + state_ids=state_ids, + position_ids=position_ids, + block_offsets=getattr(attn_metadata, 'block_offsets', None), + q_seqlens=getattr(attn_metadata, 'q_seqlens', None), + kv_seqlens=getattr(attn_metadata, 'kv_seqlens', None), + q_start_loc=getattr(attn_metadata, 'q_start_loc', None), + attn_metadata=attn_metadata, + ) + + +@dataclass +class ConceptCaches: + """ConceptLM cache views resolved once at the top-level model boundary.""" + + encoder_past_key_values: list[list[torch.Tensor]] | None = None + concept_past_key_values: list[list[torch.Tensor]] | None = None + decoder_past_key_values: list[list[torch.Tensor]] | None = None + state_caches: list[torch.Tensor] | None = None + chunk_source_idx: int = 0 + last_raw_idx: int = 1 + last_final_idx: int = 2 + + @classmethod + def build(cls, + config: PretrainedConfig, + past_key_values: list[list[torch.Tensor]] | None = None, + state_caches: list[torch.Tensor] | None = None): + """Build ConceptLM cache views from engine-provided caches.""" + encoder_past_key_values, concept_past_key_values, decoder_past_key_values = ( + _split_concept_past_key_values(config, past_key_values)) + return cls( + encoder_past_key_values=encoder_past_key_values, + concept_past_key_values=concept_past_key_values, + decoder_past_key_values=decoder_past_key_values, + state_caches=state_caches, + chunk_source_idx=int(getattr(config, 'concept_state_chunk_source_idx', 0)), + last_raw_idx=int(getattr(config, 'concept_state_last_raw_idx', 1)), + last_final_idx=int(getattr(config, 'concept_state_last_final_idx', 2)), + ) + + def state_cache(self, state_idx: int) -> torch.Tensor | None: + """Return one anonymous state-cache tensor by semantic index.""" + if self.state_caches is None: + return None + if state_idx < 0 or state_idx >= len(self.state_caches): + return None + return self.state_caches[state_idx] + + @property + def chunk_source_state(self) -> torch.Tensor | None: + """Current chunk source accumulator state cache.""" + return self.state_cache(self.chunk_source_idx) + + @property + def last_raw_states(self) -> torch.Tensor | None: + """Latest raw concept-layer state cache.""" + return self.state_cache(self.last_raw_idx) + + @property + def last_final_state(self) -> torch.Tensor | None: + """Latest final concept vector state cache.""" + return self.state_cache(self.last_final_idx) + + +@dataclass +class ConceptChunkStateUpdateResult: + """Fixed-shape result of one decode chunk-source state update.""" + + concept_input_states: torch.Tensor + next_chunk_source_states: torch.Tensor + concept_update_mask: torch.Tensor + valid_state_mask: torch.Tensor + state_ids: torch.Tensor + + def _get_configured_window(config: PretrainedConfig): """Return the reference OLMo window setting as an int or None.""" window_size = getattr(config, 'window_size', None) @@ -100,6 +213,128 @@ def _load_stacked_codebook_weight(param: torch.nn.Parameter, loaded_weight: torc target.copy_(loaded_weight) +def _split_concept_past_key_values(config: PretrainedConfig, past_key_values: list[list[torch.Tensor]] | None): + """Split the flat LMDeploy KV-cache list into ConceptLM streams.""" + if past_key_values is None or len(past_key_values) == 0: + return None, None, None + enc_layers = int(config.concept_encoder_layers) + concept_layers = int(config.concept_special_layers) + dec_layers = int(config.concept_decoder_layers) + total_layers = enc_layers + concept_layers + dec_layers + assert len(past_key_values) >= total_layers, ( + f'ConceptLM requires {total_layers} KV-cache layers ' + f'({enc_layers} encoder + {concept_layers} concept + {dec_layers} decoder), ' + f'got {len(past_key_values)}.') + enc_end = enc_layers + concept_end = enc_end + concept_layers + dec_end = concept_end + dec_layers + return past_key_values[:enc_end], past_key_values[enc_end:concept_end], past_key_values[concept_end:dec_end] + + +def _flatten_decode_position_ids(position_ids: torch.Tensor, batch_size: int) -> torch.Tensor: + """Normalize decode position ids to one absolute position per batch row.""" + if position_ids.dim() == 0: + position_ids = position_ids.view(1) + if position_ids.dim() == 1: + return position_ids.to(torch.long) + position_ids = position_ids.reshape(-1) + if position_ids.numel() == batch_size: + return position_ids.to(torch.long) + assert position_ids.numel() % batch_size == 0, ( + f'Cannot map position_ids with {position_ids.numel()} elements to batch size {batch_size}.') + return position_ids.reshape(-1, batch_size)[-1].to(torch.long) + + +def _concept_decode_chunk_state_update( + chunk_source_state_cache: torch.Tensor, + current_source_states: torch.Tensor, + state_ids: torch.Tensor, + position_ids: torch.Tensor, + chunk_size: int, + merge_method: str, +) -> ConceptChunkStateUpdateResult: + """Compute one decode step's chunk-source state update. + + This helper is intentionally fixed-shape over the decode batch. It returns + per-row next states plus a device-side boundary mask; it does not compact + concept rows by ``num_concepts_total``. The future Triton/CUDA op should + fuse this compute with the state write and skip ``state_id < 0`` rows. + + Args: + chunk_source_state_cache: ``[num_state_slots, num_sources, hidden]``. + current_source_states: ``[batch, num_sources, hidden]`` for the current + decode token after encoder source selection. These should be the + unnormalized states that are merged over the current concept chunk. + state_ids: ``[batch]`` state-cache slot per row, with ``-1`` for + padded CUDA-graph rows. + position_ids: absolute token positions for the decode rows. + chunk_size: ConceptLM chunk size. + merge_method: ``meanpooling``, ``first``, or ``last``. + """ + assert current_source_states.dim() == 3, ( + f'current_source_states must be [batch, num_sources, hidden], got {tuple(current_source_states.shape)}.') + assert chunk_source_state_cache.dim() == 3, ( + f'chunk_source_state_cache must be [num_state_slots, num_sources, hidden], ' + f'got {tuple(chunk_source_state_cache.shape)}.') + batch_size = current_source_states.size(0) + assert current_source_states.shape[1:] == chunk_source_state_cache.shape[1:], ( + f'Current source state shape {tuple(current_source_states.shape[1:])} does not match state-cache shape ' + f'{tuple(chunk_source_state_cache.shape[1:])}.') + + state_ids = state_ids.to(device=current_source_states.device, dtype=torch.long) + position_ids = _flatten_decode_position_ids(position_ids, batch_size).to(device=current_source_states.device) + assert position_ids.numel() == batch_size, ( + f'Expected {batch_size} decode position ids, got {position_ids.numel()}.') + + valid_state_mask = state_ids >= 0 + safe_state_ids = state_ids.clamp(min=0) + previous_rows = chunk_source_state_cache.index_select(0, safe_state_ids) + + chunk_size = int(chunk_size) + chunk_pos = torch.remainder(position_ids, chunk_size) + concept_update_mask = valid_state_mask & (torch.remainder(position_ids + 1, chunk_size) == 0) + first_token_mask = valid_state_mask & (chunk_pos == 0) + merge_method = str(merge_method) + + if merge_method == 'first': + update_rows = torch.where(first_token_mask.view(batch_size, 1, 1), current_source_states, previous_rows) + concept_input_states = update_rows + elif merge_method == 'last': + update_rows = current_source_states + concept_input_states = current_source_states + else: + update_rows = previous_rows + current_source_states + concept_input_states = update_rows / chunk_size + + zero_rows = torch.zeros_like(update_rows) + next_rows = torch.where(concept_update_mask.view(batch_size, 1, 1), zero_rows, update_rows) + next_rows = torch.where(valid_state_mask.view(batch_size, 1, 1), next_rows, previous_rows) + concept_input_states = torch.where(concept_update_mask.view(batch_size, 1, 1), concept_input_states, zero_rows) + return ConceptChunkStateUpdateResult( + concept_input_states=concept_input_states, + next_chunk_source_states=next_rows, + concept_update_mask=concept_update_mask, + valid_state_mask=valid_state_mask, + state_ids=state_ids, + ) + + +def _apply_concept_chunk_state_update_reference_(chunk_source_state_cache: torch.Tensor, + update: ConceptChunkStateUpdateResult): + """Reference-only state write for tests. + + This uses a Python loop and may read scalar state ids on host. Do not call + it from the serving hot path. The graph-safe implementation should write + ``update.next_chunk_source_states`` inside a backend op that skips + ``state_id < 0`` rows. + """ + for batch_idx in range(update.state_ids.numel()): + state_id = int(update.state_ids[batch_idx]) + if state_id < 0: + continue + chunk_source_state_cache[state_id].copy_(update.next_chunk_source_states[batch_idx]) + + def _qk_rmsnorm_variance(query: torch.Tensor, key: torch.Tensor) -> torch.Tensor: """Local Q/K squared sums before TP all-reduce.""" query = query.float() @@ -1266,7 +1501,40 @@ def __init__(self, self.ctx_mgr = ctx_mgr # token embedding — mirrors ``self.embedding`` in the reference. self.embedding = ConceptLMV22VQEmbedding(config, dtype=dtype, device=device) + self.encoder = ConceptLMV22VQOlmoBlock(config, + config.concept_encoder_layers, + post_layer_norm=False, + dtype=dtype, + device=device, + prefix=add_prefix('encoder', prefix)) + self.decoder = ConceptLMV22VQOlmoBlock(config, + config.concept_decoder_layers, + post_layer_norm=True, + dtype=dtype, + device=device, + prefix=add_prefix('decoder', prefix)) + self.concept_vq_input_norm = nn.LayerNorm(config.hidden_size, + eps=getattr(config, 'layernorm_epsilon', 1e-6), + dtype=dtype, + device=device) self.concept_quantizer = ConceptLMV22VQQuantizer(config, dtype=dtype, device=device) + self.concept_predictor = ConceptLMV22VQConceptPredictor(config, + dtype=dtype, + device=device, + prefix=add_prefix('concept_predictor', prefix)) + self.concept_predictor.set_attention_window(tuple(getattr(config, 'window_size', (None, None))), + getattr(config, 'window_attn_skip_freq', None)) + self.fusion_tok_norm = nn.LayerNorm(config.hidden_size, + eps=getattr(config, 'layernorm_epsilon', 1e-6), + dtype=dtype, + device=device) + self.fusion_hl_norm = nn.LayerNorm(config.hidden_size, + eps=getattr(config, 'layernorm_epsilon', 1e-6), + dtype=dtype, + device=device) + self.fusion_norm_alpha = nn.Parameter( + torch.tensor(getattr(config, 'concept_fusion_norm_alpha_init', 0.1), dtype=dtype, device=device), + requires_grad=False) self.dd_encoder_self_dd = ConceptLMV22VQSelfDD(config, config.concept_encoder_layers, use_softmax=False, @@ -1280,6 +1548,10 @@ def __init__(self, device=device) for _ in range(config.concept_decoder_layers) ]) + self.decoder_read_encoder_shared_source_norm = nn.LayerNorm(config.hidden_size, + eps=getattr(config, 'layernorm_epsilon', 1e-6), + dtype=dtype, + device=device) self.decoder_read_concept_routes = nn.ModuleList([ ConceptLMV22VQResidualRoute(config, config.concept_special_layers, @@ -1288,13 +1560,15 @@ def __init__(self, device=device) for _ in range(config.concept_decoder_layers) ]) + self.decoder_read_concept_shared_source_norm = nn.LayerNorm(config.hidden_size, + eps=getattr(config, 'layernorm_epsilon', 1e-6), + dtype=dtype, + device=device) + self.final_read_concept_gate_logits = nn.Parameter( + torch.zeros(config.concept_decoder_layers, 2, dtype=dtype, device=device), + requires_grad=False) self.dd_two_route_add = ConceptLMV22VQTwoRouteAdd(config, dtype=dtype, device=device) - self.concept_predictor = ConceptLMV22VQConceptPredictor(config, - dtype=dtype, - device=device, - prefix=add_prefix('concept_predictor', prefix)) - self.concept_predictor.set_attention_window(tuple(getattr(config, 'window_size', (None, None))), - getattr(config, 'window_attn_skip_freq', None)) + self.concept_ops = ConceptLMRuntimeOps() # output projection — mirrors ``output_layer`` in the reference. Built # via build_lm_head and named ``lm_head`` so DeployModelMixinV1. # get_logits picks it up directly; load_weights maps the checkpoint's @@ -1308,8 +1582,13 @@ def forward(self, past_key_values: list[list[torch.Tensor]], attn_metadata=None, inputs_embeds: torch.Tensor = None, + state_ids: torch.Tensor | None = None, + state_caches: list[torch.Tensor] | None = None, **kwargs): """Model forward, return hidden_states (logits computed by runtime).""" + concept_metadata = self._build_concept_metadata(position_ids, attn_metadata, state_ids) + concept_caches = self._build_concept_caches(past_key_values, state_caches) + _ = concept_metadata, concept_caches if inputs_embeds is None: # NOTE: placeholder. The real path is # embed -> encoder -> concept(vq+predictor) -> fusion -> decoder @@ -1323,6 +1602,219 @@ def get_input_embeddings(self): """Get input embeddings.""" return self.embedding.word_embeddings + def get_output_embeddings(self): + """Get output embeddings.""" + return self.lm_head + + def _split_past_key_values(self, past_key_values: list[list[torch.Tensor]] | None): + """Split the flat LMDeploy KV-cache list into ConceptLM streams.""" + return _split_concept_past_key_values(self.config, past_key_values) + + def _build_concept_metadata(self, + position_ids: torch.Tensor | None, + attn_metadata: Any = None, + state_ids: torch.Tensor | None = None) -> ConceptMetadata: + """Build top-level ConceptLM metadata for future serving paths.""" + return ConceptMetadata.build( + self.config, + position_ids=position_ids, + attn_metadata=attn_metadata, + state_ids=state_ids, + ) + + def _build_concept_caches(self, + past_key_values: list[list[torch.Tensor]] | None, + state_caches: list[torch.Tensor] | None = None) -> ConceptCaches: + """Build top-level ConceptLM cache views for future serving paths.""" + return ConceptCaches.build( + self.config, + past_key_values=past_key_values, + state_caches=state_caches, + ) + + def _decode_chunk_state_update(self, + current_source_states: torch.Tensor, + concept_metadata: ConceptMetadata, + concept_caches: ConceptCaches) -> ConceptChunkStateUpdateResult: + """Update decode chunk-source state and return fixed-shape rows. + + CUDA uses the Triton writer. CPU uses the reference writer for tests. + The returned rows deliberately avoid dynamic concept-row compaction, + matching the CUDA graph route in the design doc. + """ + chunk_source_state = concept_caches.chunk_source_state + if chunk_source_state is None: + raise RuntimeError('ConceptLM decode chunk update requires concept chunk source state cache.') + if concept_metadata.state_ids is None: + raise RuntimeError('ConceptLM decode chunk update requires state_ids.') + if concept_metadata.position_ids is None: + raise RuntimeError('ConceptLM decode chunk update requires position_ids.') + concept_input_states, next_rows, update_mask = self.concept_ops.decode_chunk_state_update( + chunk_source_state, + current_source_states, + concept_metadata.state_ids, + concept_metadata.position_ids, + concept_metadata.chunk_size, + concept_metadata.merge_method, + ) + state_ids = concept_metadata.state_ids.to(device=current_source_states.device, dtype=torch.long) + return ConceptChunkStateUpdateResult( + concept_input_states=concept_input_states, + next_chunk_source_states=next_rows, + concept_update_mask=update_mask, + valid_state_mask=state_ids >= 0, + state_ids=state_ids, + ) + + def support_cuda_graph( + self, + input_ids: torch.Tensor, + position_ids: torch.Tensor, + past_key_values: list[list[torch.Tensor]], + attn_metadata: Any = None, + inputs_embeds: torch.Tensor = None, + **kwargs, + ): + """Disable CUDA graph until ConceptLM decode state update is graph-safe. + + ``states_shapes`` makes the engine allocate graph-padded state ids. The + current top-level ConceptLM helpers still include dense/reference-only + chunk operations, so allowing the default decode graph capture would + bake in the wrong execution contract. + """ + return False + + def _merge_chunks(self, hidden_states: torch.Tensor) -> torch.Tensor: + """Reference chunk merge for dense ``[seq, batch, hidden]`` states.""" + assert hidden_states.dim() == 3, ( + f'_merge_chunks currently expects dense [seq, batch, hidden], got {tuple(hidden_states.shape)}.') + seq_len, batch_size, hidden_size = hidden_states.shape + chunk_size = int(self.config.concept_chunk_size) + if seq_len < chunk_size: + return hidden_states.mean(dim=0, keepdim=True) + + usable = (seq_len // chunk_size) * chunk_size + chunks = hidden_states[:usable].reshape(usable // chunk_size, chunk_size, batch_size, hidden_size) + merge_method = getattr(self.config, 'concept_chunk_merge_method', 'meanpooling') + if merge_method == 'first': + return chunks[:, 0] + if merge_method == 'last': + return chunks[:, -1] + return chunks.mean(dim=1) + + def _repeat_shift(self, concept_states: torch.Tensor, seq_len: int) -> torch.Tensor: + """Repeat chunk-level concept states back to the token timeline.""" + chunk_size = int(self.config.concept_chunk_size) + shifted = torch.cat((torch.zeros_like(concept_states[:1]), concept_states), dim=0) + repeated = shifted.repeat_interleave(chunk_size, dim=0)[:int(seq_len) + 1] + if bool(getattr(self.config, 'concept_shift_feature', True)): + repeated = repeated[1:] + else: + repeated = repeated[:seq_len] + if repeated.shape[0] < seq_len: + pad_len = int(seq_len) - repeated.shape[0] + repeated = torch.cat((repeated, repeated.new_zeros(pad_len, *repeated.shape[1:])), dim=0) + return repeated[:seq_len] + + def _build_encoder_concept_states(self, encoder_raw_states: list[torch.Tensor]) -> torch.Tensor: + """Build chunk-level encoder states used by the concept predictor.""" + chunks = [self._merge_chunks(state) for state in encoder_raw_states[:-1]] + assert len(chunks) > 0, 'ConceptLM concept-read-encoder route requires at least one encoder source state.' + states = torch.stack(chunks, dim=2) + return self.concept_predictor.normalize_encoder_concept_states(states) + + def _route_gate(self, layer_idx: int) -> torch.Tensor: + """Return decoder route gate ``[decoder_dd_scale, concept_route_scale]``.""" + return self.final_read_concept_gate_logits[int(layer_idx)].float().softmax(dim=-1) + + def _encode(self, + hidden_states: torch.Tensor, + position_ids: torch.Tensor, + past_key_values: list[list[torch.Tensor]] | None = None, + attn_metadata: Any = None): + """Encoder stack plus encoder self-DD. + + This helper mirrors the reference flow but consumes LMDeploy attention + inputs. It is wired for the future full forward path; continuous + batching still needs caller-side concept metadata before the full model + can use it safely. + """ + raw_states = [] + history_buffer = self.dd_encoder_self_dd.make_history_buffer(hidden_states) + self.dd_encoder_self_dd.write_history(history_buffer, 0, hidden_states) + rotary_pos_emb = self.encoder._make_rotary_pos_emb(hidden_states, position_ids) + + for layer_idx, layer in enumerate(self.encoder.layers): + pkv = past_key_values[layer_idx] if past_key_values is not None else None + raw = layer( + hidden_states, + rotary_pos_emb=rotary_pos_emb, + past_key_value=pkv, + attn_metadata=attn_metadata, + ) + raw_states.append(raw) + self.dd_encoder_self_dd.write_history(history_buffer, layer_idx + 1, raw) + hidden_states = self.dd_encoder_self_dd.forward_from_buffer(layer_idx, raw, history_buffer) + return hidden_states, raw_states + + def _decode(self, + decoder_input: torch.Tensor, + encoder_raw_states: list[torch.Tensor], + final_concept_state: torch.Tensor, + concept_raw_states: list[torch.Tensor], + position_ids: torch.Tensor, + past_key_values: list[list[torch.Tensor]] | None = None, + attn_metadata: Any = None): + """Decoder stack plus decoder DD and residual routes.""" + decoder_encoder_states = torch.stack(tuple(encoder_raw_states), dim=2) + decoder_encoder_states = self.decoder_read_encoder_shared_source_norm(decoder_encoder_states) + + concept_states = torch.stack(tuple(concept_raw_states), dim=2) + zero_chunk = torch.zeros_like(concept_states[:1]) + concept_states = torch.cat((zero_chunk, concept_states), dim=0) + concept_states = self.decoder_read_concept_shared_source_norm(concept_states) + + chunk_size = int(self.config.concept_chunk_size) + hidden_states = decoder_input + history_buffer = self.dd_two_route_add.make_history_buffer(hidden_states) + self.dd_two_route_add.write_history(history_buffer, 0, hidden_states) + rotary_pos_emb = self.decoder._make_rotary_pos_emb(hidden_states, position_ids) + + for layer_idx, layer in enumerate(self.decoder.layers): + pkv = past_key_values[layer_idx] if past_key_values is not None else None + raw = layer( + hidden_states, + rotary_pos_emb=rotary_pos_emb, + past_key_value=pkv, + attn_metadata=attn_metadata, + ) + self.dd_two_route_add.write_history(history_buffer, layer_idx + 1, raw) + gate = self._route_gate(layer_idx) + hidden_states = self.dd_two_route_add.forward_from_buffer( + layer_idx, + raw, + history_buffer, + final_concept_state, + gate[0], + ) + hidden_states = self.decoder_read_encoder_routes[layer_idx]( + hidden_states, + decoder_encoder_states, + source_dim=2, + ) + hidden_states = self.decoder_read_concept_routes[layer_idx].forward_repeated_chunks( + hidden_states, + concept_states, + chunk_size, + bool(getattr(self.config, 'concept_shift_feature', True)), + residual_scale=gate[1], + source_dim=2, + ) + + if self.decoder.final_layernorm is not None: + hidden_states = self.decoder.final_layernorm(hidden_states) + return hidden_states + def prepare_inputs_for_generation(self, past_key_values: list[list[torch.Tensor]], inputs_embeds: torch.Tensor | None = None, @@ -1337,32 +1829,35 @@ def prepare_inputs_for_generation(self, past_key_values=past_key_values, attn_metadata=attn_metadata, inputs_embeds=inputs_embeds, + state_ids=context.state_offsets, + state_caches=context.state_caches, ) def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]): - """Load weights. - - Only modules currently wired into the top-level model are loaded. Other - checkpoint tensors are skipped until their corresponding modules are - added. - """ + """Load native ConceptLM checkpoint weights into implemented modules.""" # (checkpoint_name, target_name) weight_map = { 'embedding.word_embeddings.weight': 'embedding.word_embeddings.weight', 'output_layer.weight': 'lm_head.weight', } codebook_prefix = 'concept_quantizer.codebook.' - concept_hlm_prefix = 'concept_predictor.hlm_block.' prediction_head_prefix = 'concept_predictor.prediction_heads.' + block_prefixes = ( + ('encoder.', self.encoder), + ('decoder.', self.decoder), + ('concept_predictor.hlm_block.', self.concept_predictor.hlm_block), + ) params_dict = dict(self.named_parameters()) for name, loaded_weight in weights: if 'rotary_emb.inv_freq' in name: continue - if name.startswith(concept_hlm_prefix): - self.concept_predictor.hlm_block.load_weights( - [(name, loaded_weight)], - prefix=concept_hlm_prefix[:-1], - ) + loaded_by_block = False + for block_prefix, block in block_prefixes: + if name.startswith(block_prefix): + block.load_weights([(name, loaded_weight)], prefix=block_prefix[:-1]) + loaded_by_block = True + break + if loaded_by_block: continue if name.startswith(prediction_head_prefix): suffix = name[len(prediction_head_prefix):] @@ -1388,7 +1883,7 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]): if target is None: target = name if target not in params_dict: - # skip modules not yet wired into the top-level model + # skip checkpoint metadata or tensors owned by future runtime-only paths continue param = params_dict[target] load_weight(param, loaded_weight) diff --git a/lmdeploy/pytorch/nn/__init__.py b/lmdeploy/pytorch/nn/__init__.py index eae4c4324b..7f0a7c53a6 100644 --- a/lmdeploy/pytorch/nn/__init__.py +++ b/lmdeploy/pytorch/nn/__init__.py @@ -3,6 +3,7 @@ # https://github.com/vllm-project/vllm/blob/main/vllm/attention/ from .activation import GeluAndMul, SiluAndMul # noqa: F401 from .attention import Attention, FlashAttention # noqa: F401 +from .conceptlm import ConceptLMRuntimeOps # noqa: F401 from .embedding import ParallelEmbedding # noqa: F401 from .hc_prepost import HcPrePost # noqa: F401 from .norm import LayerNorm, RMSNorm, rms_scale # noqa: F401 diff --git a/lmdeploy/pytorch/nn/conceptlm.py b/lmdeploy/pytorch/nn/conceptlm.py new file mode 100644 index 0000000000..1aeb2a505f --- /dev/null +++ b/lmdeploy/pytorch/nn/conceptlm.py @@ -0,0 +1,56 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from torch import Tensor, nn + +from lmdeploy.pytorch.backends import OpType, get_backend + + +class ConceptLMRuntimeOps(nn.Module): + """ConceptLM model-specific runtime operation wrapper. + + The model calls this nn module only. Backend implementations own dispatch, + and CUDA implementations own direct Triton kernel launchers. + """ + + def __init__(self): + super().__init__() + backend = get_backend() + builder = backend.get_layer_impl_builder(OpType.ConceptLMRuntimeOps) + self.impl = builder.build() + + def decode_chunk_state_update( + self, + chunk_source_state_cache: Tensor, + current_source_states: Tensor, + state_ids: Tensor, + position_ids: Tensor, + chunk_size: int, + merge_method: str, + ) -> tuple[Tensor, Tensor, Tensor]: + """Update state cache and return concept inputs, next rows, and mask.""" + return self.impl.decode_chunk_state_update( + chunk_source_state_cache, + current_source_states, + state_ids, + position_ids, + chunk_size, + merge_method, + ) + + def forward( + self, + chunk_source_state_cache: Tensor, + current_source_states: Tensor, + state_ids: Tensor, + position_ids: Tensor, + chunk_size: int, + merge_method: str, + ) -> tuple[Tensor, Tensor, Tensor]: + """Alias the current runtime op for module-call compatibility.""" + return self.decode_chunk_state_update( + chunk_source_state_cache, + current_source_states, + state_ids, + position_ids, + chunk_size, + merge_method, + ) From ed3887eef4897801125d9f80ece96ecefadb18a0 Mon Sep 17 00:00:00 2001 From: grimoire Date: Fri, 24 Jul 2026 18:31:23 +0800 Subject: [PATCH 03/16] Wire ConceptLM packed prefill forward --- lmdeploy/pytorch/models/intern_ncp.py | 483 ++++++++++++++++++++++---- 1 file changed, 425 insertions(+), 58 deletions(-) diff --git a/lmdeploy/pytorch/models/intern_ncp.py b/lmdeploy/pytorch/models/intern_ncp.py index 41c439e3c4..67db3f74da 100644 --- a/lmdeploy/pytorch/models/intern_ncp.py +++ b/lmdeploy/pytorch/models/intern_ncp.py @@ -16,14 +16,16 @@ - ``_ConceptPredictor``: concept block container and prediction heads. - top-level encoder/decoder containers, fusion norms, route norms, and checkpoint loading for the implemented module tree. + - packed non-decode prefill path through encoder -> concept + predictor/quantizer -> fusion -> decoder, with per-request chunk-stream + attention metadata derived explicitly from token-stream metadata. -Results are NOT yet correct end-to-end — the encoder/concept/decoder control -flow still needs runtime metadata for the compressed concept stream. The -module structure mirrors the reference for readability. +Decode still needs the graph-safe compressed concept-stream runtime contract. +The module structure mirrors the reference for readability. """ from collections.abc import Iterable -from dataclasses import dataclass +from dataclasses import dataclass, replace from typing import Any import torch @@ -156,6 +158,56 @@ class ConceptChunkStateUpdateResult: state_ids: torch.Tensor +@dataclass +class ConceptPrefillMetadata: + """Packed prefill metadata derived once from token attention metadata. + + Field groups: + - token stream: original engine-provided token request boundaries. + - concept stream: compact chunk-token request boundaries and positions + used by the concept predictor attention. + - chunk merge: token -> concept ids and helper ids used to reduce encoder + token states into concept states without a Python batch loop. + - repeat/gather: concept -> token ids used to project compact concept + states back to the packed token stream. + - scalar bounds: eager compact sizes / upper bounds needed by metadata and + attention launch parameters. + """ + + # Token stream metadata, shape [batch]. This is the original packed prefill + # layout consumed by normal token attention. + token_q_seqlens: torch.Tensor + token_q_start_loc: torch.Tensor + + # Concept stream metadata, shape [batch] plus compact concept positions. + # These describe the shorter chunk-token stream consumed by concept + # predictor attention. + concept_q_seqlens: torch.Tensor + concept_q_start_loc: torch.Tensor + concept_position_ids: torch.Tensor + + # Chunk merge metadata. ``merge_token_to_concept`` maps each packed token to + # the compact concept row that owns it, or -1 when the token is dropped from + # concept production. Counts/first/last ids implement mean/first/last merge. + merge_token_to_concept: torch.Tensor + merge_token_counts: torch.Tensor + merge_first_token_ids: torch.Tensor + merge_last_token_ids: torch.Tensor + merge_short_concept_mask: torch.Tensor + + # Repeat/gather metadata. Maps each packed token row to the compact concept + # row it should read after shift semantics are applied, or -1 for the + # zero-concept row. + token_to_concept: torch.Tensor + + # Scalar sizes/bounds. ``num_concepts_total`` is the exact compact size in + # eager prefill; ``max_concepts_per_request`` is the per-request attention + # launch bound. + num_tokens_total: int + num_concepts_total: int + max_concepts_per_request: int + + def _get_configured_window(config: PretrainedConfig): """Return the reference OLMo window setting as an int or None.""" window_size = getattr(config, 'window_size', None) @@ -1588,15 +1640,28 @@ def forward(self, """Model forward, return hidden_states (logits computed by runtime).""" concept_metadata = self._build_concept_metadata(position_ids, attn_metadata, state_ids) concept_caches = self._build_concept_caches(past_key_values, state_caches) - _ = concept_metadata, concept_caches if inputs_embeds is None: - # NOTE: placeholder. The real path is - # embed -> encoder -> concept(vq+predictor) -> fusion -> decoder - # added incrementally. Returns raw embeddings as hidden_states. hidden_states = self.embedding(input_ids) else: hidden_states = inputs_embeds - return hidden_states + + if concept_metadata.is_decoding: + # TODO: wire the decode control flow using ConceptLMRuntimeOps + # state-cache updates. Keeping the previous placeholder behavior + # avoids mixing decode-state work into the prefill patch. + return hidden_states + + hidden_states, prefill_position_ids = self._normalize_prefill_inputs( + hidden_states, + position_ids, + attn_metadata, + ) + return self._forward_prefill_packed( + hidden_states, + prefill_position_ids, + concept_metadata, + concept_caches, + ) def get_input_embeddings(self): """Get input embeddings.""" @@ -1684,48 +1749,328 @@ def support_cuda_graph( """ return False - def _merge_chunks(self, hidden_states: torch.Tensor) -> torch.Tensor: - """Reference chunk merge for dense ``[seq, batch, hidden]`` states.""" - assert hidden_states.dim() == 3, ( - f'_merge_chunks currently expects dense [seq, batch, hidden], got {tuple(hidden_states.shape)}.') - seq_len, batch_size, hidden_size = hidden_states.shape - chunk_size = int(self.config.concept_chunk_size) + def _route_gate(self, layer_idx: int) -> torch.Tensor: + """Return decoder route gate ``[decoder_dd_scale, concept_route_scale]``.""" + return self.final_read_concept_gate_logits[int(layer_idx)].float().softmax(dim=-1) + + def _normalize_prefill_inputs(self, + hidden_states: torch.Tensor, + position_ids: torch.Tensor, + attn_metadata: Any = None): + """Normalize LMDeploy prefill inputs to packed token layout. + + Engine layout stays fixed as ``[1, total_tokens, hidden]`` with + ``position_ids=[1, total_tokens]``. Per-request boundaries come from + ``attn_metadata.q_seqlens``/``q_start_loc`` and are used only to build + the concept stream. + """ + if attn_metadata is None or getattr(attn_metadata, 'q_seqlens', None) is None: + raise RuntimeError('ConceptLM prefill requires q_seqlens attention metadata.') + if getattr(attn_metadata, 'is_decoding', False): + raise RuntimeError('ConceptLM prefill input normalization received decode metadata.') + + if hidden_states.dim() != 3 or hidden_states.size(0) != 1: + raise NotImplementedError( + f'ConceptLM prefill expects fixed engine layout [1, total_tokens, hidden], ' + f'got {tuple(hidden_states.shape)}.') + + total_tokens = hidden_states.size(1) + if position_ids.dim() != 2 or position_ids.size(0) != 1: + raise NotImplementedError( + f'ConceptLM prefill expects fixed position layout [1, total_tokens], ' + f'got {tuple(position_ids.shape)}.') + if position_ids.size(1) != total_tokens: + raise ValueError(f'position_ids length {position_ids.size(1)} does not match token length {total_tokens}.') + + q_seqlens = attn_metadata.q_seqlens + if q_seqlens.dim() != 1: + raise ValueError(f'ConceptLM prefill expects 1-D q_seqlens, got {tuple(q_seqlens.shape)}.') + + hidden_states = hidden_states[0].contiguous() + return hidden_states, position_ids[0].to(device=hidden_states.device) + + @staticmethod + def _concept_count_from_seq_len(seq_len: int, chunk_size: int) -> int: + """Return reference ConceptLM chunk count for one request length.""" + seq_len = int(seq_len or 0) + if seq_len <= 0: + return 0 if seq_len < chunk_size: - return hidden_states.mean(dim=0, keepdim=True) + return 1 + return seq_len // chunk_size + + @staticmethod + def _concept_counts_from_q_seqlens(q_seqlens: torch.Tensor, chunk_size: int) -> torch.Tensor: + """Vectorized version of ``_concept_count_from_seq_len``.""" + counts = torch.div(q_seqlens, chunk_size, rounding_mode='floor').clamp(min=1) + return torch.where(q_seqlens > 0, counts, torch.zeros_like(q_seqlens)) + + @staticmethod + def _repeat_slot_ids(token_pos: torch.Tensor, chunk_size: int, shift_feature: bool) -> torch.Tensor: + """Return local concept slot read by each token after shift semantics.""" + if shift_feature: + return torch.div(token_pos + 1, chunk_size, rounding_mode='floor') - 1 + return torch.div(token_pos, chunk_size, rounding_mode='floor') - 1 + + @staticmethod + def _get_max_concepts_per_request(token_attn_metadata: Any, + concept_q_seqlens: torch.Tensor, + chunk_size: int) -> int: + """Return per-request concept attention bound without hidden context access.""" + max_q_seqlen = getattr(token_attn_metadata, 'max_q_seqlen', None) + if max_q_seqlen is not None: + return ConceptLMV22VQForCausalLM._concept_count_from_seq_len(int(max_q_seqlen), chunk_size) + + # Test/direct-call fallback. Serving should use the scheduler-provided + # Python max_q_seqlen above, as in DSV4 metadata construction. + return int(concept_q_seqlens.max().item()) if concept_q_seqlens.numel() > 0 else 0 + + def _build_prefill_metadata(self, + token_attn_metadata: Any, + position_ids: torch.Tensor) -> ConceptPrefillMetadata: + """Build packed token-to-concept metadata for batched prefill.""" + if token_attn_metadata is None: + raise RuntimeError('ConceptLM prefill requires attention metadata.') + if getattr(token_attn_metadata, 'is_decoding', False): + raise RuntimeError('ConceptLM prefill metadata cannot be built from decode metadata.') + + q_seqlens = token_attn_metadata.q_seqlens + q_start_loc = getattr(token_attn_metadata, 'q_start_loc', None) + if q_start_loc is None: + q_start_loc = F.pad(torch.cumsum(q_seqlens, dim=0, dtype=torch.int32), (1, 0))[:-1] + + total_tokens = int(position_ids.numel()) + chunk_size = int(self.config.concept_chunk_size) + shift_feature = bool(getattr(self.config, 'concept_shift_feature', True)) + + q_seqlens_long = q_seqlens.to(dtype=torch.long, device=position_ids.device) + q_start_loc_long = q_start_loc.to(dtype=torch.long, device=position_ids.device) + cu_q_seqlens = getattr(token_attn_metadata, 'cu_seqlens_q', None) + if cu_q_seqlens is None: + cu_q_seqlens = F.pad(torch.cumsum(q_seqlens, dim=0, dtype=torch.int32), (1, 0)) + cu_q_seqlens_long = cu_q_seqlens.to(dtype=torch.long, device=position_ids.device) + + token_ids = torch.arange(total_tokens, dtype=torch.long, device=position_ids.device) + token_seq = torch.searchsorted(cu_q_seqlens_long[1:], token_ids, right=True) + token_pos = token_ids - cu_q_seqlens_long[token_seq] + + concept_q_seqlens_long = self._concept_counts_from_q_seqlens(q_seqlens_long, chunk_size) + concept_q_seqlens = concept_q_seqlens_long.to(dtype=q_seqlens.dtype, device=q_seqlens.device) + concept_q_start_loc = F.pad(torch.cumsum(concept_q_seqlens, dim=0, dtype=torch.int32), (1, 0))[:-1] + concept_q_start_loc_long = concept_q_start_loc.to(dtype=torch.long, device=position_ids.device) + concept_cu_seqlens_long = F.pad(torch.cumsum(concept_q_seqlens_long, dim=0), (1, 0)) + + # TODO: remove this eager compact-size scalar read by moving ConceptLM + # concept-stream allocation/compaction into the engine/backend contract, + # like DSV4's precomputed metadata and Qwen3.5's state-cache metadata. + num_concepts_total = int(concept_cu_seqlens_long[-1].item()) + concept_ids = torch.arange(num_concepts_total, dtype=torch.long, device=position_ids.device) + concept_seq = torch.searchsorted(concept_cu_seqlens_long[1:], concept_ids, right=True) + local_concept_ids = concept_ids - concept_cu_seqlens_long[concept_seq] + concept_token_start = q_start_loc_long[concept_seq] + local_concept_ids * chunk_size + concept_position_ids = position_ids[concept_token_start] + + seq_concept_start = concept_q_start_loc_long[token_seq] + seq_concept_count = concept_q_seqlens_long[token_seq] + repeat_slots = self._repeat_slot_ids(token_pos, chunk_size, shift_feature) + valid_repeat = (repeat_slots >= 0) & (repeat_slots < seq_concept_count) + token_to_concept = torch.where( + valid_repeat, + seq_concept_start + repeat_slots, + torch.full_like(repeat_slots, -1), + ) + + merge_slots = torch.div(token_pos, chunk_size, rounding_mode='floor') + valid_merge = (merge_slots >= 0) & (merge_slots < seq_concept_count) + merge_token_to_concept = torch.where( + valid_merge, + seq_concept_start + merge_slots, + torch.full_like(merge_slots, -1), + ) + safe_merge_ids = merge_token_to_concept.clamp(min=0) + merge_token_counts = torch.zeros(num_concepts_total, dtype=torch.int32, device=position_ids.device) + merge_token_counts.index_add_(0, safe_merge_ids, valid_merge.to(dtype=torch.int32)) + + concept_seq_len = q_seqlens_long[concept_seq] + merge_first_pos = torch.where( + concept_seq_len < chunk_size, + torch.zeros_like(local_concept_ids), + local_concept_ids * chunk_size, + ) + merge_last_pos = torch.where( + concept_seq_len < chunk_size, + (concept_seq_len - 1).clamp(min=0), + local_concept_ids * chunk_size + chunk_size - 1, + ) + merge_first_token_ids = q_start_loc_long[concept_seq] + merge_first_pos + merge_last_token_ids = q_start_loc_long[concept_seq] + merge_last_pos + merge_short_concept_mask = concept_seq_len < chunk_size + max_concepts_per_request = self._get_max_concepts_per_request( + token_attn_metadata, + concept_q_seqlens_long, + chunk_size, + ) - usable = (seq_len // chunk_size) * chunk_size - chunks = hidden_states[:usable].reshape(usable // chunk_size, chunk_size, batch_size, hidden_size) + return ConceptPrefillMetadata( + token_q_seqlens=q_seqlens, + token_q_start_loc=q_start_loc, + concept_q_seqlens=concept_q_seqlens, + concept_q_start_loc=concept_q_start_loc, + concept_position_ids=concept_position_ids, + merge_token_to_concept=merge_token_to_concept, + merge_token_counts=merge_token_counts, + merge_first_token_ids=merge_first_token_ids, + merge_last_token_ids=merge_last_token_ids, + merge_short_concept_mask=merge_short_concept_mask, + token_to_concept=token_to_concept, + num_tokens_total=total_tokens, + num_concepts_total=num_concepts_total, + max_concepts_per_request=max_concepts_per_request, + ) + + @staticmethod + def _merge_chunks_mean_packed(hidden_states: torch.Tensor, + prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: + """Mean-pool packed token states by precomputed concept ids.""" + merge_token_to_concept = prefill_metadata.merge_token_to_concept.to(device=hidden_states.device) + valid_merge = (merge_token_to_concept >= 0).to(dtype=hidden_states.dtype).unsqueeze(-1) + safe_merge_ids = merge_token_to_concept.clamp(min=0) + merged = hidden_states.new_zeros((prefill_metadata.num_concepts_total, hidden_states.size(-1))) + merged.index_add_(0, safe_merge_ids, hidden_states * valid_merge) + counts = prefill_metadata.merge_token_counts.clamp(min=1).to(device=hidden_states.device, + dtype=hidden_states.dtype) + return merged / counts.unsqueeze(-1) + + def _merge_chunks_packed(self, + hidden_states: torch.Tensor, + prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: + """Merge packed token states into packed per-request concept states.""" + assert hidden_states.dim() == 2, ( + f'_merge_chunks_packed expects [total_tokens, hidden], got {tuple(hidden_states.shape)}.') merge_method = getattr(self.config, 'concept_chunk_merge_method', 'meanpooling') + if prefill_metadata.num_concepts_total <= 0: + raise RuntimeError('ConceptLM prefill produced no concept chunks.') if merge_method == 'first': - return chunks[:, 0] + merged = hidden_states[prefill_metadata.merge_first_token_ids] + short_mean = self._merge_chunks_mean_packed(hidden_states, prefill_metadata) + return torch.where(prefill_metadata.merge_short_concept_mask[:, None], short_mean, merged) if merge_method == 'last': - return chunks[:, -1] - return chunks.mean(dim=1) + merged = hidden_states[prefill_metadata.merge_last_token_ids] + short_mean = self._merge_chunks_mean_packed(hidden_states, prefill_metadata) + return torch.where(prefill_metadata.merge_short_concept_mask[:, None], short_mean, merged) - def _repeat_shift(self, concept_states: torch.Tensor, seq_len: int) -> torch.Tensor: - """Repeat chunk-level concept states back to the token timeline.""" - chunk_size = int(self.config.concept_chunk_size) - shifted = torch.cat((torch.zeros_like(concept_states[:1]), concept_states), dim=0) - repeated = shifted.repeat_interleave(chunk_size, dim=0)[:int(seq_len) + 1] - if bool(getattr(self.config, 'concept_shift_feature', True)): - repeated = repeated[1:] - else: - repeated = repeated[:seq_len] - if repeated.shape[0] < seq_len: - pad_len = int(seq_len) - repeated.shape[0] - repeated = torch.cat((repeated, repeated.new_zeros(pad_len, *repeated.shape[1:])), dim=0) - return repeated[:seq_len] - - def _build_encoder_concept_states(self, encoder_raw_states: list[torch.Tensor]) -> torch.Tensor: - """Build chunk-level encoder states used by the concept predictor.""" - chunks = [self._merge_chunks(state) for state in encoder_raw_states[:-1]] + return self._merge_chunks_mean_packed(hidden_states, prefill_metadata) + + @staticmethod + def _gather_zero_prefixed_concepts(concept_states_with_zero: torch.Tensor, + prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: + """Gather zero-prefixed concept rows to packed token rows.""" + token_to_concept = prefill_metadata.token_to_concept.to(device=concept_states_with_zero.device) + gather_ids = torch.clamp(token_to_concept + 1, min=0) + return concept_states_with_zero[gather_ids] + + def _repeat_shift_packed(self, + concept_states: torch.Tensor, + prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: + """Gather packed concept states back to packed token states.""" + concept_states_with_zero = torch.cat((torch.zeros_like(concept_states[:1]), concept_states), dim=0) + return self._gather_zero_prefixed_concepts(concept_states_with_zero, prefill_metadata) + + def _repeat_shift_source_states_packed(self, + concept_states_with_zero: torch.Tensor, + prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: + """Gather zero-prefixed packed concept source states to token rows.""" + return self._gather_zero_prefixed_concepts(concept_states_with_zero, prefill_metadata) + + def _build_encoder_concept_states_packed(self, + encoder_raw_states: list[torch.Tensor], + prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: + """Build packed chunk-level encoder states used by the concept predictor.""" + chunks = [self._merge_chunks_packed(state, prefill_metadata) for state in encoder_raw_states[:-1]] assert len(chunks) > 0, 'ConceptLM concept-read-encoder route requires at least one encoder source state.' - states = torch.stack(chunks, dim=2) + states = torch.stack(chunks, dim=-2) return self.concept_predictor.normalize_encoder_concept_states(states) - def _route_gate(self, layer_idx: int) -> torch.Tensor: - """Return decoder route gate ``[decoder_dd_scale, concept_route_scale]``.""" - return self.final_read_concept_gate_logits[int(layer_idx)].float().softmax(dim=-1) + def _build_concept_prefill_metadata(self, + token_attn_metadata: Any, + prefill_metadata: ConceptPrefillMetadata): + """Build chunk-stream attention metadata for packed prefill.""" + if token_attn_metadata is None: + raise RuntimeError('ConceptLM prefill requires attention metadata for concept predictor attention.') + if getattr(token_attn_metadata, 'is_decoding', False): + raise RuntimeError('ConceptLM concept prefill metadata cannot be built from decode metadata.') + + concept_q_seqlens = prefill_metadata.concept_q_seqlens + concept_q_start_loc = prefill_metadata.concept_q_start_loc + concept_cu_seqlens = torch.nn.functional.pad( + torch.cumsum(concept_q_seqlens, dim=0, dtype=torch.int32), + (1, 0), + ) + max_concept_seqlen = int(prefill_metadata.max_concepts_per_request) + + updates = dict( + is_decoding=False, + q_start_loc=concept_q_start_loc, + q_seqlens=concept_q_seqlens, + kv_seqlens=concept_q_seqlens, + cu_seqlens_q=concept_cu_seqlens, + cu_seqlens_k=concept_cu_seqlens, + ) + if hasattr(token_attn_metadata, 'kv_start_loc'): + updates['kv_start_loc'] = concept_q_start_loc + if hasattr(token_attn_metadata, 'kv_flatten_size'): + updates['kv_flatten_size'] = int(prefill_metadata.num_concepts_total) + if hasattr(token_attn_metadata, 'max_q_seqlen'): + updates['max_q_seqlen'] = max_concept_seqlen + if hasattr(token_attn_metadata, 'max_kv_seqlen'): + updates['max_kv_seqlen'] = max_concept_seqlen + return replace(token_attn_metadata, **updates) + + def _forward_prefill_packed(self, + hidden_states: torch.Tensor, + position_ids: torch.Tensor, + concept_metadata: ConceptMetadata, + concept_caches: ConceptCaches): + """Packed non-decode ConceptLM forward.""" + if (concept_caches.encoder_past_key_values is None or concept_caches.concept_past_key_values is None + or concept_caches.decoder_past_key_values is None): + raise RuntimeError('ConceptLM prefill requires encoder, concept, and decoder KV caches.') + prefill_metadata = self._build_prefill_metadata(concept_metadata.attn_metadata, position_ids) + hidden_states, encoder_raw_states = self._encode( + hidden_states, + position_ids, + past_key_values=concept_caches.encoder_past_key_values, + attn_metadata=concept_metadata.attn_metadata, + ) + concept_hidden = self._merge_chunks_packed(hidden_states, prefill_metadata) + concept_hidden = self.concept_vq_input_norm(concept_hidden) + encoder_concept_states = self._build_encoder_concept_states_packed(encoder_raw_states, prefill_metadata) + concept_attn_metadata = self._build_concept_prefill_metadata( + concept_metadata.attn_metadata, + prefill_metadata, + ) + concept_logits, concept_raw_states = self.concept_predictor( + concept_hidden, + encoder_concept_states, + prefill_metadata.concept_position_ids, + past_key_values=concept_caches.concept_past_key_values, + attn_metadata=concept_attn_metadata, + ) + predicted_vectors = self.concept_quantizer(concept_logits) + repeated_concepts = self._repeat_shift_packed(predicted_vectors, prefill_metadata) + decoder_input = self.fusion_tok_norm(hidden_states) + self.fusion_norm_alpha.to( + hidden_states.dtype) * self.fusion_hl_norm(repeated_concepts.to(hidden_states.dtype)) + final_hidden = self._decode( + decoder_input, + encoder_raw_states, + repeated_concepts, + concept_raw_states, + position_ids, + past_key_values=concept_caches.decoder_past_key_values, + attn_metadata=concept_metadata.attn_metadata, + prefill_metadata=prefill_metadata, + ) + return final_hidden.unsqueeze(0).contiguous() def _encode(self, hidden_states: torch.Tensor, @@ -1764,15 +2109,29 @@ def _decode(self, concept_raw_states: list[torch.Tensor], position_ids: torch.Tensor, past_key_values: list[list[torch.Tensor]] | None = None, - attn_metadata: Any = None): + attn_metadata: Any = None, + prefill_metadata: ConceptPrefillMetadata | None = None): """Decoder stack plus decoder DD and residual routes.""" - decoder_encoder_states = torch.stack(tuple(encoder_raw_states), dim=2) - decoder_encoder_states = self.decoder_read_encoder_shared_source_norm(decoder_encoder_states) + if prefill_metadata is None: + decoder_encoder_states = torch.stack(tuple(encoder_raw_states), dim=2) + decoder_encoder_states = self.decoder_read_encoder_shared_source_norm(decoder_encoder_states) + + concept_states = torch.stack(tuple(concept_raw_states), dim=2) + zero_chunk = torch.zeros_like(concept_states[:1]) + concept_states = torch.cat((zero_chunk, concept_states), dim=0) + concept_states = self.decoder_read_concept_shared_source_norm(concept_states) + decoder_encoder_source_dim = 2 + repeated_concept_states = None + else: + decoder_encoder_states = torch.stack(tuple(encoder_raw_states), dim=-2) + decoder_encoder_states = self.decoder_read_encoder_shared_source_norm(decoder_encoder_states) - concept_states = torch.stack(tuple(concept_raw_states), dim=2) - zero_chunk = torch.zeros_like(concept_states[:1]) - concept_states = torch.cat((zero_chunk, concept_states), dim=0) - concept_states = self.decoder_read_concept_shared_source_norm(concept_states) + concept_states = torch.stack(tuple(concept_raw_states), dim=-2) + zero_chunk = torch.zeros_like(concept_states[:1]) + concept_states = torch.cat((zero_chunk, concept_states), dim=0) + concept_states = self.decoder_read_concept_shared_source_norm(concept_states) + decoder_encoder_source_dim = -2 + repeated_concept_states = self._repeat_shift_source_states_packed(concept_states, prefill_metadata) chunk_size = int(self.config.concept_chunk_size) hidden_states = decoder_input @@ -1800,16 +2159,24 @@ def _decode(self, hidden_states = self.decoder_read_encoder_routes[layer_idx]( hidden_states, decoder_encoder_states, - source_dim=2, - ) - hidden_states = self.decoder_read_concept_routes[layer_idx].forward_repeated_chunks( - hidden_states, - concept_states, - chunk_size, - bool(getattr(self.config, 'concept_shift_feature', True)), - residual_scale=gate[1], - source_dim=2, + source_dim=decoder_encoder_source_dim, ) + if repeated_concept_states is None: + hidden_states = self.decoder_read_concept_routes[layer_idx].forward_repeated_chunks( + hidden_states, + concept_states, + chunk_size, + bool(getattr(self.config, 'concept_shift_feature', True)), + residual_scale=gate[1], + source_dim=2, + ) + else: + hidden_states = self.decoder_read_concept_routes[layer_idx]( + hidden_states, + repeated_concept_states, + residual_scale=gate[1], + source_dim=-2, + ) if self.decoder.final_layernorm is not None: hidden_states = self.decoder.final_layernorm(hidden_states) From 98b61f6835771cfb8eb79ef179b1669822884dce Mon Sep 17 00:00:00 2001 From: grimoire Date: Fri, 24 Jul 2026 19:14:19 +0800 Subject: [PATCH 04/16] Wire ConceptLM eager decode path --- lmdeploy/pytorch/configurations/conceptlm.py | 7 +- lmdeploy/pytorch/models/intern_ncp.py | 422 ++++++++++++++++++- 2 files changed, 417 insertions(+), 12 deletions(-) diff --git a/lmdeploy/pytorch/configurations/conceptlm.py b/lmdeploy/pytorch/configurations/conceptlm.py index c4e5273728..c66b269ac9 100644 --- a/lmdeploy/pytorch/configurations/conceptlm.py +++ b/lmdeploy/pytorch/configurations/conceptlm.py @@ -72,8 +72,13 @@ def build(cls, hf_config, model_path: str = None, **kwargs): hidden_size = int(hf_config.hidden_size) state_dtype = _get_concept_state_dtype(hf_config) concept_encoder_read_sources = max(enc_layers - 1, 0) + # Decode accumulates the current chunk for every state needed when a + # chunk boundary emits one concept. Row 0 is the final encoder hidden + # used as concept-predictor input; following rows are encoder raw states + # consumed by concept-read-encoder residual routes. + concept_chunk_state_sources = 1 + concept_encoder_read_sources model_config.states_shapes = [ - ((concept_encoder_read_sources, hidden_size), state_dtype), + ((concept_chunk_state_sources, hidden_size), state_dtype), ((concept_layers, hidden_size), state_dtype), ((hidden_size, ), state_dtype), ] diff --git a/lmdeploy/pytorch/models/intern_ncp.py b/lmdeploy/pytorch/models/intern_ncp.py index 67db3f74da..f8121a515d 100644 --- a/lmdeploy/pytorch/models/intern_ncp.py +++ b/lmdeploy/pytorch/models/intern_ncp.py @@ -156,6 +156,22 @@ class ConceptChunkStateUpdateResult: concept_update_mask: torch.Tensor valid_state_mask: torch.Tensor state_ids: torch.Tensor + safe_state_ids: torch.Tensor + + +@dataclass +class ConceptDecodeMetadata: + """Fixed-layout decode metadata derived once from engine inputs. + + Decode is always represented as the engine's fixed ``[1, batch]`` token + layout at the model boundary and flattened to ``[batch]`` / ``[batch, H]`` + only inside ConceptLM helpers. + """ + + position_ids: torch.Tensor + state_ids: torch.Tensor + safe_state_ids: torch.Tensor + valid_state_mask: torch.Tensor @dataclass @@ -368,6 +384,7 @@ def _concept_decode_chunk_state_update( concept_update_mask=concept_update_mask, valid_state_mask=valid_state_mask, state_ids=state_ids, + safe_state_ids=safe_state_ids, ) @@ -1646,10 +1663,12 @@ def forward(self, hidden_states = inputs_embeds if concept_metadata.is_decoding: - # TODO: wire the decode control flow using ConceptLMRuntimeOps - # state-cache updates. Keeping the previous placeholder behavior - # avoids mixing decode-state work into the prefill patch. - return hidden_states + return self._forward_decode( + hidden_states, + position_ids, + concept_metadata, + concept_caches, + ) hidden_states, prefill_position_ids = self._normalize_prefill_inputs( hidden_states, @@ -1729,6 +1748,7 @@ def _decode_chunk_state_update(self, concept_update_mask=update_mask, valid_state_mask=state_ids >= 0, state_ids=state_ids, + safe_state_ids=state_ids.clamp(min=0), ) def support_cuda_graph( @@ -1740,12 +1760,13 @@ def support_cuda_graph( inputs_embeds: torch.Tensor = None, **kwargs, ): - """Disable CUDA graph until ConceptLM decode state update is graph-safe. + """Disable CUDA graph until ConceptLM concept-boundary updates are graph-safe. ``states_shapes`` makes the engine allocate graph-padded state ids. The - current top-level ConceptLM helpers still include dense/reference-only - chunk operations, so allowing the default decode graph capture would - bake in the wrong execution contract. + eager decode path updates chunk-source state through a backend op, but + still compacts boundary concept rows dynamically before concept + predictor attention. Capturing that would bake in a batch-specific + concept update shape. """ return False @@ -1789,6 +1810,290 @@ def _normalize_prefill_inputs(self, hidden_states = hidden_states[0].contiguous() return hidden_states, position_ids[0].to(device=hidden_states.device) + def _normalize_decode_inputs(self, hidden_states: torch.Tensor, position_ids: torch.Tensor): + """Normalize LMDeploy decode inputs to one flat row per active request.""" + if hidden_states.dim() != 3 or hidden_states.size(0) != 1: + raise NotImplementedError( + f'ConceptLM decode expects fixed engine layout [1, batch, hidden], ' + f'got {tuple(hidden_states.shape)}.') + batch_size = hidden_states.size(1) + position_ids = _flatten_decode_position_ids(position_ids, batch_size).to(device=hidden_states.device) + if position_ids.numel() != batch_size: + raise ValueError(f'Expected {batch_size} decode position ids, got {position_ids.numel()}.') + return hidden_states[0].contiguous(), position_ids + + @staticmethod + def _build_decode_metadata(position_ids: torch.Tensor, + state_ids: torch.Tensor | None, + batch_size: int, + device: torch.device) -> ConceptDecodeMetadata: + """Build fixed-shape decode metadata from engine state ids.""" + if state_ids is None: + raise RuntimeError('ConceptLM decode requires state_ids.') + state_ids = state_ids.to(device=device, dtype=torch.long).reshape(-1) + if state_ids.numel() != batch_size: + raise ValueError(f'Expected {batch_size} decode state ids, got {state_ids.numel()}.') + valid_state_mask = state_ids >= 0 + return ConceptDecodeMetadata( + position_ids=position_ids, + state_ids=state_ids, + safe_state_ids=state_ids.clamp(min=0), + valid_state_mask=valid_state_mask, + ) + + @staticmethod + def _select_decode_state_rows(state_cache: torch.Tensor, + decode_metadata: ConceptDecodeMetadata) -> torch.Tensor: + """Gather state-cache rows and zero out padded decode rows.""" + rows = state_cache.index_select(0, decode_metadata.safe_state_ids) + mask_shape = (decode_metadata.valid_state_mask.size(0), ) + (1, ) * (rows.dim() - 1) + valid_mask = decode_metadata.valid_state_mask.view(mask_shape) + return torch.where(valid_mask, rows, torch.zeros_like(rows)) + + def _decode_concept_read_mask(self, decode_metadata: ConceptDecodeMetadata) -> torch.Tensor: + """Return rows whose current decode token should read a cached concept.""" + repeat_slots = self._repeat_slot_ids( + decode_metadata.position_ids, + int(self.config.concept_chunk_size), + bool(getattr(self.config, 'concept_shift_feature', True)), + ) + return decode_metadata.valid_state_mask & (repeat_slots >= 0) + + def _build_decode_chunk_source_states(self, + hidden_states: torch.Tensor, + encoder_raw_states: list[torch.Tensor]) -> torch.Tensor: + """Build current per-row states accumulated until the next concept boundary. + + Row 0 is the final encoder hidden, used as concept-predictor input when + a boundary is reached. Remaining rows mirror the prefill + ``encoder_raw_states[:-1]`` route sources. + """ + source_states = [hidden_states] + source_states.extend(encoder_raw_states[:-1]) + return torch.stack(tuple(source_states), dim=1) + + @staticmethod + def _decode_concept_position_ids(position_ids: torch.Tensor, chunk_size: int) -> torch.Tensor: + """Return reference RoPE positions for concept rows emitted at decode boundaries.""" + return (position_ids - int(chunk_size) + 1).clamp(min=0) + + def _build_concept_decode_metadata_eager(self, + token_attn_metadata: Any, + decode_metadata: ConceptDecodeMetadata, + boundary_indices: torch.Tensor): + """Build dynamic concept-stream decode metadata for boundary rows. + + This is the eager-only bridge: it compacts rows whose token completed a + concept chunk and makes concept KV positions advance on the compressed + concept timeline. CUDA graph support should replace this with a + fixed-shape backend metadata object/op instead of calling ``nonzero`` + and constructing variable-batch attention metadata here. + """ + if token_attn_metadata is None or getattr(token_attn_metadata, 'block_offsets', None) is None: + raise RuntimeError('ConceptLM decode concept update requires token attention metadata.') + + device = decode_metadata.position_ids.device + num_boundary_rows = boundary_indices.numel() + q_seqlens = getattr(token_attn_metadata, 'q_seqlens', None) + q_start_loc = getattr(token_attn_metadata, 'q_start_loc', None) + kv_seqlens = getattr(token_attn_metadata, 'kv_seqlens', None) + if q_seqlens is None or q_start_loc is None or kv_seqlens is None: + raise RuntimeError('ConceptLM decode concept update requires q/q_start/kv sequence metadata.') + q_dtype = q_seqlens.dtype + q_start_dtype = q_start_loc.dtype + kv_dtype = kv_seqlens.dtype + + concept_q_seqlens = torch.ones((num_boundary_rows, ), dtype=q_dtype, device=device) + concept_q_start_loc = torch.arange(num_boundary_rows, dtype=q_start_dtype, device=device) + concept_cu_seqlens = F.pad(torch.cumsum(concept_q_seqlens, dim=0, dtype=torch.int32), (1, 0)) + concept_kv_seqlens = torch.div( + decode_metadata.position_ids.index_select(0, boundary_indices) + 1, + int(self.config.concept_chunk_size), + rounding_mode='floor', + ).to(dtype=kv_dtype) + + updates = dict( + is_decoding=True, + block_offsets=token_attn_metadata.block_offsets.index_select(0, boundary_indices), + q_start_loc=concept_q_start_loc, + q_seqlens=concept_q_seqlens, + kv_seqlens=concept_kv_seqlens, + cu_seqlens_q=concept_cu_seqlens, + cu_seqlens_k=concept_cu_seqlens, + ) + if hasattr(token_attn_metadata, 'kv_start_loc'): + updates['kv_start_loc'] = concept_kv_seqlens - concept_q_seqlens.to(dtype=concept_kv_seqlens.dtype) + if hasattr(token_attn_metadata, 'kv_flatten_size'): + updates['kv_flatten_size'] = num_boundary_rows + if hasattr(token_attn_metadata, 'max_q_seqlen'): + updates['max_q_seqlen'] = 1 + if hasattr(token_attn_metadata, 'max_kv_seqlen'): + # Decode kernels use kv_seqlens directly. Keep a conservative bound + # without reading the dynamic maximum back to host. + updates['max_kv_seqlen'] = getattr(token_attn_metadata, 'max_kv_seqlen') + if hasattr(token_attn_metadata, 'scheduler_metadata'): + updates['scheduler_metadata'] = None + if hasattr(token_attn_metadata, 'tile_scheduler_metadata'): + updates['tile_scheduler_metadata'] = None + if hasattr(token_attn_metadata, 'num_splits'): + updates['num_splits'] = None + if hasattr(token_attn_metadata, 'fill_seqlens'): + updates['fill_seqlens'] = concept_q_seqlens + return replace(token_attn_metadata, **updates) + + @staticmethod + def _stack_concept_raw_states(concept_raw_states: list[torch.Tensor]) -> torch.Tensor: + """Stack raw concept-layer states to ``[rows, concept_layers, hidden]``.""" + return torch.stack(tuple(concept_raw_states), dim=1) + + def _write_decode_concept_states_eager_(self, + concept_caches: ConceptCaches, + decode_metadata: ConceptDecodeMetadata, + boundary_indices: torch.Tensor, + predicted_vectors: torch.Tensor, + concept_raw_states: list[torch.Tensor]): + """Write newly emitted concept states to persistent decode caches.""" + last_final_state = concept_caches.last_final_state + last_raw_states = concept_caches.last_raw_states + if last_final_state is None or last_raw_states is None: + raise RuntimeError('ConceptLM decode concept update requires last concept state caches.') + state_ids = decode_metadata.safe_state_ids.index_select(0, boundary_indices) + last_final_state.index_copy_(0, state_ids, predicted_vectors.to(dtype=last_final_state.dtype)) + raw_rows = self._stack_concept_raw_states(concept_raw_states) + last_raw_states.index_copy_(0, state_ids, raw_rows.to(dtype=last_raw_states.dtype)) + + def _update_decode_concept_states_eager_(self, + chunk_update: ConceptChunkStateUpdateResult, + decode_metadata: ConceptDecodeMetadata, + concept_metadata: ConceptMetadata, + concept_caches: ConceptCaches): + """Emit and cache concept states for decode rows that complete a chunk. + + TODO: replace this dynamic eager bridge with a graph-safe backend op and + fixed-shape concept attention metadata. The model-level flow should stay + the same: chunk accumulator -> concept predictor on boundary rows -> + cached final/raw concept states -> decoder routes. + """ + boundary_indices = torch.nonzero(chunk_update.concept_update_mask, as_tuple=False).flatten() + if boundary_indices.numel() == 0: + return + if concept_caches.concept_past_key_values is None: + raise RuntimeError('ConceptLM decode concept update requires concept KV caches.') + + boundary_concept_inputs = chunk_update.concept_input_states.index_select(0, boundary_indices) + concept_hidden = self.concept_vq_input_norm(boundary_concept_inputs[:, 0]) + encoder_concept_states = self.concept_predictor.normalize_encoder_concept_states( + boundary_concept_inputs[:, 1:]) + concept_position_ids = self._decode_concept_position_ids( + decode_metadata.position_ids, + concept_metadata.chunk_size, + ).index_select(0, boundary_indices) + concept_attn_metadata = self._build_concept_decode_metadata_eager( + concept_metadata.attn_metadata, + decode_metadata, + boundary_indices, + ) + concept_logits, concept_raw_states = self.concept_predictor( + concept_hidden, + encoder_concept_states, + concept_position_ids, + past_key_values=concept_caches.concept_past_key_values, + attn_metadata=concept_attn_metadata, + ) + predicted_vectors = self.concept_quantizer(concept_logits) + self._write_decode_concept_states_eager_( + concept_caches, + decode_metadata, + boundary_indices, + predicted_vectors, + concept_raw_states, + ) + + def _merge_prefill_tail_chunk_states(self, + source_states: torch.Tensor, + prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: + """Build per-request partial chunk accumulator rows after prefill.""" + assert source_states.dim() == 3, ( + f'Expected source states [total_tokens, num_sources, hidden], got {tuple(source_states.shape)}.') + device = source_states.device + chunk_size = int(self.config.concept_chunk_size) + q_seqlens = prefill_metadata.token_q_seqlens.to(device=device, dtype=torch.long) + q_start_loc = prefill_metadata.token_q_start_loc.to(device=device, dtype=torch.long) + batch_size = q_seqlens.size(0) + tail_lens = torch.remainder(q_seqlens, chunk_size) + tail_lens = torch.where(q_seqlens < chunk_size, q_seqlens, tail_lens) + tail_lens = torch.where(q_seqlens > 0, tail_lens, torch.zeros_like(tail_lens)) + has_tail = tail_lens > 0 + + tail_rows = source_states.new_zeros((batch_size, source_states.size(1), source_states.size(2))) + if source_states.size(0) == 0: + return tail_rows + merge_method = getattr(self.config, 'concept_chunk_merge_method', 'meanpooling') + if merge_method == 'first': + token_ids = q_start_loc + q_seqlens - tail_lens + token_ids = token_ids.clamp(min=0, max=max(source_states.size(0) - 1, 0)) + rows = source_states[token_ids] + return torch.where(has_tail.view(batch_size, 1, 1), rows, tail_rows) + if merge_method == 'last': + token_ids = q_start_loc + q_seqlens - 1 + token_ids = token_ids.clamp(min=0, max=max(source_states.size(0) - 1, 0)) + rows = source_states[token_ids] + return torch.where(has_tail.view(batch_size, 1, 1), rows, tail_rows) + + token_ids = torch.arange(prefill_metadata.num_tokens_total, dtype=torch.long, device=device) + cu_q_seqlens = F.pad(torch.cumsum(q_seqlens, dim=0), (1, 0)) + token_seq = torch.searchsorted(cu_q_seqlens[1:], token_ids, right=True) + token_pos = token_ids - cu_q_seqlens[token_seq] + token_tail_start = q_seqlens[token_seq] - tail_lens[token_seq] + valid_tail = (tail_lens[token_seq] > 0) & (token_pos >= token_tail_start) + weighted_source = source_states * valid_tail.to(dtype=source_states.dtype).view(-1, 1, 1) + tail_rows.index_add_(0, token_seq, weighted_source) + return tail_rows + + def _write_prefill_state_caches_eager_(self, + concept_caches: ConceptCaches, + concept_metadata: ConceptMetadata, + prefill_metadata: ConceptPrefillMetadata, + source_states: torch.Tensor, + predicted_vectors: torch.Tensor, + concept_raw_states: list[torch.Tensor]): + """Seed decode state caches from a completed prefill forward.""" + if concept_caches.state_caches is None or concept_metadata.state_ids is None: + return + chunk_source_state = concept_caches.chunk_source_state + last_raw_states = concept_caches.last_raw_states + last_final_state = concept_caches.last_final_state + if chunk_source_state is None or last_raw_states is None or last_final_state is None: + raise RuntimeError('ConceptLM prefill state init requires all ConceptLM state caches.') + + state_ids = concept_metadata.state_ids.to(device=source_states.device, dtype=torch.long).reshape(-1) + batch_size = prefill_metadata.token_q_seqlens.numel() + if state_ids.numel() != batch_size: + raise ValueError(f'Expected {batch_size} prefill state ids, got {state_ids.numel()}.') + valid_indices = torch.nonzero(state_ids >= 0, as_tuple=False).flatten() + if valid_indices.numel() == 0: + return + + valid_state_ids = state_ids.index_select(0, valid_indices) + tail_rows = self._merge_prefill_tail_chunk_states(source_states, prefill_metadata) + chunk_source_state.index_copy_(0, valid_state_ids, + tail_rows.index_select(0, valid_indices).to(dtype=chunk_source_state.dtype)) + + concept_counts = prefill_metadata.concept_q_seqlens.to(device=source_states.device, dtype=torch.long) + concept_start = prefill_metadata.concept_q_start_loc.to(device=source_states.device, dtype=torch.long) + concept_valid_mask = (state_ids >= 0) & (concept_counts > 0) + concept_indices = torch.nonzero(concept_valid_mask, as_tuple=False).flatten() + if concept_indices.numel() == 0: + return + + concept_state_ids = state_ids.index_select(0, concept_indices) + last_concept_ids = concept_start.index_select(0, concept_indices) + concept_counts.index_select( + 0, concept_indices) - 1 + last_final_rows = predicted_vectors.index_select(0, last_concept_ids) + last_final_state.index_copy_(0, concept_state_ids, last_final_rows.to(dtype=last_final_state.dtype)) + raw_rows = self._stack_concept_raw_states(concept_raw_states).index_select(0, last_concept_ids) + last_raw_states.index_copy_(0, concept_state_ids, raw_rows.to(dtype=last_raw_states.dtype)) + @staticmethod def _concept_count_from_seq_len(seq_len: int, chunk_size: int) -> int: """Return reference ConceptLM chunk count for one request length.""" @@ -2057,6 +2362,14 @@ def _forward_prefill_packed(self, attn_metadata=concept_attn_metadata, ) predicted_vectors = self.concept_quantizer(concept_logits) + self._write_prefill_state_caches_eager_( + concept_caches, + concept_metadata, + prefill_metadata, + self._build_decode_chunk_source_states(hidden_states, encoder_raw_states), + predicted_vectors, + concept_raw_states, + ) repeated_concepts = self._repeat_shift_packed(predicted_vectors, prefill_metadata) decoder_input = self.fusion_tok_norm(hidden_states) + self.fusion_norm_alpha.to( hidden_states.dtype) * self.fusion_hl_norm(repeated_concepts.to(hidden_states.dtype)) @@ -2072,6 +2385,85 @@ def _forward_prefill_packed(self, ) return final_hidden.unsqueeze(0).contiguous() + def _forward_decode(self, + hidden_states: torch.Tensor, + position_ids: torch.Tensor, + concept_metadata: ConceptMetadata, + concept_caches: ConceptCaches): + """Eager ConceptLM decode path. + + This path is semantically structured for serving but intentionally keeps + CUDA graph disabled. Boundary concept updates currently use dynamic row + compaction in ``_update_decode_concept_states_eager_``; the stable + graph route needs a fixed-shape concept metadata/backend op. + """ + if concept_caches.encoder_past_key_values is None or concept_caches.decoder_past_key_values is None: + raise RuntimeError('ConceptLM decode requires encoder and decoder KV caches.') + if concept_caches.chunk_source_state is None: + raise RuntimeError('ConceptLM decode requires chunk source state cache.') + if concept_caches.last_raw_states is None or concept_caches.last_final_state is None: + raise RuntimeError('ConceptLM decode requires cached last concept states.') + + hidden_states, decode_position_ids = self._normalize_decode_inputs(hidden_states, position_ids) + decode_metadata = self._build_decode_metadata( + decode_position_ids, + concept_metadata.state_ids, + hidden_states.size(0), + hidden_states.device, + ) + decode_concept_metadata = replace( + concept_metadata, + position_ids=decode_position_ids, + state_ids=decode_metadata.state_ids, + ) + + hidden_states, encoder_raw_states = self._encode( + hidden_states, + decode_position_ids, + past_key_values=concept_caches.encoder_past_key_values, + attn_metadata=concept_metadata.attn_metadata, + ) + previous_final_concept_state = self._select_decode_state_rows(concept_caches.last_final_state, decode_metadata) + previous_concept_raw_state_rows = self._select_decode_state_rows(concept_caches.last_raw_states, + decode_metadata) + current_source_states = self._build_decode_chunk_source_states(hidden_states, encoder_raw_states) + chunk_update = self._decode_chunk_state_update( + current_source_states, + decode_concept_metadata, + concept_caches, + ) + self._update_decode_concept_states_eager_( + chunk_update, + decode_metadata, + decode_concept_metadata, + concept_caches, + ) + + if bool(getattr(self.config, 'concept_shift_feature', True)): + final_concept_state = self._select_decode_state_rows(concept_caches.last_final_state, decode_metadata) + concept_raw_state_rows = self._select_decode_state_rows(concept_caches.last_raw_states, decode_metadata) + else: + final_concept_state = previous_final_concept_state + concept_raw_state_rows = previous_concept_raw_state_rows + concept_read_mask = self._decode_concept_read_mask(decode_metadata) + final_concept_state = torch.where(concept_read_mask.view(-1, 1), final_concept_state, + torch.zeros_like(final_concept_state)) + concept_raw_state_rows = torch.where(concept_read_mask.view(-1, 1, 1), concept_raw_state_rows, + torch.zeros_like(concept_raw_state_rows)) + decoder_input = self.fusion_tok_norm(hidden_states) + self.fusion_norm_alpha.to( + hidden_states.dtype) * self.fusion_hl_norm(final_concept_state.to(hidden_states.dtype)) + final_hidden = self._decode( + decoder_input, + encoder_raw_states, + final_concept_state, + concept_raw_states=[], + position_ids=decode_position_ids, + past_key_values=concept_caches.decoder_past_key_values, + attn_metadata=concept_metadata.attn_metadata, + decode_concept_states=concept_raw_state_rows, + ) + return final_hidden.unsqueeze(0).contiguous() + def _encode(self, hidden_states: torch.Tensor, position_ids: torch.Tensor, @@ -2106,13 +2498,21 @@ def _decode(self, decoder_input: torch.Tensor, encoder_raw_states: list[torch.Tensor], final_concept_state: torch.Tensor, - concept_raw_states: list[torch.Tensor], + concept_raw_states: list[torch.Tensor] | None, position_ids: torch.Tensor, past_key_values: list[list[torch.Tensor]] | None = None, attn_metadata: Any = None, - prefill_metadata: ConceptPrefillMetadata | None = None): + prefill_metadata: ConceptPrefillMetadata | None = None, + decode_concept_states: torch.Tensor | None = None): """Decoder stack plus decoder DD and residual routes.""" - if prefill_metadata is None: + if decode_concept_states is not None: + decoder_encoder_states = torch.stack(tuple(encoder_raw_states), dim=-2) + decoder_encoder_states = self.decoder_read_encoder_shared_source_norm(decoder_encoder_states) + + concept_states = self.decoder_read_concept_shared_source_norm(decode_concept_states) + decoder_encoder_source_dim = -2 + repeated_concept_states = concept_states + elif prefill_metadata is None: decoder_encoder_states = torch.stack(tuple(encoder_raw_states), dim=2) decoder_encoder_states = self.decoder_read_encoder_shared_source_norm(decoder_encoder_states) From 786b758f2b2437a16a02c0357edddfe8134d7159 Mon Sep 17 00:00:00 2001 From: grimoire Date: Fri, 24 Jul 2026 20:42:11 +0800 Subject: [PATCH 05/16] Enable ConceptLM decode CUDA graph path --- lmdeploy/pytorch/backends/conceptlm.py | 46 ++- lmdeploy/pytorch/backends/cuda/conceptlm.py | 49 ++- .../pytorch/backends/default/conceptlm.py | 61 ++++ lmdeploy/pytorch/kernels/cuda/conceptlm.py | 295 ++++++++++++++++++ lmdeploy/pytorch/models/intern_ncp.py | 293 +++++++++-------- lmdeploy/pytorch/nn/conceptlm.py | 62 +++- 6 files changed, 672 insertions(+), 134 deletions(-) diff --git a/lmdeploy/pytorch/backends/conceptlm.py b/lmdeploy/pytorch/backends/conceptlm.py index a02974c707..5aec30e84d 100644 --- a/lmdeploy/pytorch/backends/conceptlm.py +++ b/lmdeploy/pytorch/backends/conceptlm.py @@ -7,9 +7,8 @@ class ConceptLMRuntimeOpsImpl(ABC): """ConceptLM runtime operation implementation. - Model-specific runtime/cache operations live behind this single backend - interface. That keeps model code free from direct kernel calls while - avoiding one OpType/nn module per small ConceptLM state operation. + Model-specific runtime/cache operations live behind this single backend interface. That keeps model code free from + direct kernel calls while avoiding one OpType/nn module per small ConceptLM state operation. """ @abstractmethod @@ -22,7 +21,46 @@ def decode_chunk_state_update( chunk_size: int, merge_method: str, ) -> tuple[Tensor, Tensor, Tensor]: - """Update state cache and return concept inputs, next rows, and mask.""" + """Update state cache and return concept inputs, next rows, and + mask.""" + raise NotImplementedError('Not implemented.') + + @abstractmethod + def decode_kv_cache_snapshot( + self, + k_cache: Tensor, + v_cache: Tensor, + block_offsets: Tensor, + kv_seqlens: Tensor, + ) -> tuple[Tensor, Tensor]: + """Snapshot one decode KV slot per batch row.""" + raise NotImplementedError('Not implemented.') + + @abstractmethod + def decode_kv_cache_restore( + self, + k_cache: Tensor, + v_cache: Tensor, + saved_k: Tensor, + saved_v: Tensor, + block_offsets: Tensor, + kv_seqlens: Tensor, + restore_mask: Tensor, + ) -> None: + """Restore one decode KV slot for masked batch rows.""" + raise NotImplementedError('Not implemented.') + + @abstractmethod + def decode_concept_state_update( + self, + last_raw_state_cache: Tensor, + last_final_state_cache: Tensor, + predicted_vectors: Tensor, + raw_states: Tensor, + state_ids: Tensor, + update_mask: Tensor, + ) -> None: + """Write final/raw concept states for masked decode rows.""" raise NotImplementedError('Not implemented.') diff --git a/lmdeploy/pytorch/backends/cuda/conceptlm.py b/lmdeploy/pytorch/backends/cuda/conceptlm.py index 951e34ac96..61672957f2 100644 --- a/lmdeploy/pytorch/backends/cuda/conceptlm.py +++ b/lmdeploy/pytorch/backends/cuda/conceptlm.py @@ -1,7 +1,12 @@ # Copyright (c) OpenMMLab. All rights reserved. from torch import Tensor -from lmdeploy.pytorch.kernels.cuda.conceptlm import decode_chunk_state_update +from lmdeploy.pytorch.kernels.cuda.conceptlm import ( + decode_chunk_state_update, + decode_concept_state_update, + decode_kv_cache_restore, + decode_kv_cache_snapshot, +) from ..conceptlm import ConceptLMRuntimeOpsBuilder, ConceptLMRuntimeOpsImpl @@ -28,6 +33,48 @@ def decode_chunk_state_update( merge_method, ) + def decode_kv_cache_snapshot( + self, + k_cache: Tensor, + v_cache: Tensor, + block_offsets: Tensor, + kv_seqlens: Tensor, + ) -> tuple[Tensor, Tensor]: + """Snapshot one decode KV slot per batch row.""" + return decode_kv_cache_snapshot(k_cache, v_cache, block_offsets, kv_seqlens) + + def decode_kv_cache_restore( + self, + k_cache: Tensor, + v_cache: Tensor, + saved_k: Tensor, + saved_v: Tensor, + block_offsets: Tensor, + kv_seqlens: Tensor, + restore_mask: Tensor, + ) -> None: + """Restore one decode KV slot for masked batch rows.""" + return decode_kv_cache_restore(k_cache, v_cache, saved_k, saved_v, block_offsets, kv_seqlens, restore_mask) + + def decode_concept_state_update( + self, + last_raw_state_cache: Tensor, + last_final_state_cache: Tensor, + predicted_vectors: Tensor, + raw_states: Tensor, + state_ids: Tensor, + update_mask: Tensor, + ) -> None: + """Write final/raw concept states for masked decode rows.""" + return decode_concept_state_update( + last_raw_state_cache, + last_final_state_cache, + predicted_vectors, + raw_states, + state_ids, + update_mask, + ) + class TritonConceptLMRuntimeOpsBuilder(ConceptLMRuntimeOpsBuilder): """Triton ConceptLM runtime operation builder.""" diff --git a/lmdeploy/pytorch/backends/default/conceptlm.py b/lmdeploy/pytorch/backends/default/conceptlm.py index 558dc4f0f0..ada11b9866 100644 --- a/lmdeploy/pytorch/backends/default/conceptlm.py +++ b/lmdeploy/pytorch/backends/default/conceptlm.py @@ -22,6 +22,19 @@ def _flatten_decode_position_ids(position_ids: Tensor, batch_size: int, device: class DefaultConceptLMRuntimeOpsImpl(ConceptLMRuntimeOpsImpl): """Torch fallback implementation of ConceptLM runtime operations.""" + @staticmethod + def _decode_kv_cache_rows(k_cache: Tensor, + block_offsets: Tensor, + kv_seqlens: Tensor) -> tuple[Tensor, Tensor]: + """Return cache block ids and page offsets for one slot per row.""" + block_size = k_cache.size(1) + kv_seqlens = kv_seqlens.to(device=block_offsets.device, dtype=torch.long).clamp(min=1) + slot_ids = kv_seqlens - 1 + block_idx = torch.div(slot_ids, block_size, rounding_mode='floor') + page_offsets = torch.remainder(slot_ids, block_size) + block_ids = block_offsets.to(dtype=torch.long).gather(1, block_idx.view(-1, 1)).view(-1) + return block_ids, page_offsets + def decode_chunk_state_update( self, chunk_source_state_cache: Tensor, @@ -78,6 +91,54 @@ def decode_chunk_state_update( chunk_source_state_cache[state_id].copy_(next_rows[batch_idx]) return concept_input_states, next_rows, update_mask + def decode_kv_cache_snapshot( + self, + k_cache: Tensor, + v_cache: Tensor, + block_offsets: Tensor, + kv_seqlens: Tensor, + ) -> tuple[Tensor, Tensor]: + """Snapshot one decode KV slot per batch row.""" + block_ids, page_offsets = self._decode_kv_cache_rows(k_cache, block_offsets, kv_seqlens) + return k_cache[block_ids, page_offsets].clone(), v_cache[block_ids, page_offsets].clone() + + def decode_kv_cache_restore( + self, + k_cache: Tensor, + v_cache: Tensor, + saved_k: Tensor, + saved_v: Tensor, + block_offsets: Tensor, + kv_seqlens: Tensor, + restore_mask: Tensor, + ) -> None: + """Restore one decode KV slot for masked batch rows.""" + block_ids, page_offsets = self._decode_kv_cache_rows(k_cache, block_offsets, kv_seqlens) + restore_mask = restore_mask.to(device=k_cache.device, dtype=torch.bool).view(-1, 1, 1) + current_k = k_cache[block_ids, page_offsets] + current_v = v_cache[block_ids, page_offsets] + k_cache[block_ids, page_offsets] = torch.where(restore_mask, saved_k, current_k) + v_cache[block_ids, page_offsets] = torch.where(restore_mask, saved_v, current_v) + + def decode_concept_state_update( + self, + last_raw_state_cache: Tensor, + last_final_state_cache: Tensor, + predicted_vectors: Tensor, + raw_states: Tensor, + state_ids: Tensor, + update_mask: Tensor, + ) -> None: + """Write final/raw concept states for masked decode rows.""" + state_ids = state_ids.to(device=predicted_vectors.device, dtype=torch.long).reshape(-1) + update_mask = update_mask.to(device=predicted_vectors.device, dtype=torch.bool).reshape(-1) + for batch_idx in range(state_ids.numel()): + state_id = int(state_ids[batch_idx]) + if state_id < 0 or not bool(update_mask[batch_idx]): + continue + last_final_state_cache[state_id].copy_(predicted_vectors[batch_idx].to(last_final_state_cache.dtype)) + last_raw_state_cache[state_id].copy_(raw_states[batch_idx].to(last_raw_state_cache.dtype)) + class DefaultConceptLMRuntimeOpsBuilder(ConceptLMRuntimeOpsBuilder): """Torch fallback ConceptLM runtime operation builder.""" diff --git a/lmdeploy/pytorch/kernels/cuda/conceptlm.py b/lmdeploy/pytorch/kernels/cuda/conceptlm.py index aaee94c7db..fb53178dbb 100644 --- a/lmdeploy/pytorch/kernels/cuda/conceptlm.py +++ b/lmdeploy/pytorch/kernels/cuda/conceptlm.py @@ -81,6 +81,167 @@ def _decode_chunk_state_update_kernel( tl.store(update_mask + batch_id, valid_state & is_boundary) +@triton.jit +def _decode_kv_cache_snapshot_kernel( + k_cache, + v_cache, + block_offsets, + kv_seqlens, + saved_k, + saved_v, + k_stride_n, + k_stride_b, + k_stride_h, + k_stride_d, + v_stride_n, + v_stride_b, + v_stride_h, + v_stride_d, + boff_stride_b, + boff_stride_n, + sk_stride_b, + sk_stride_h, + sk_stride_d, + sv_stride_b, + sv_stride_h, + sv_stride_d, + HEAD_DIM: tl.constexpr, + TOTAL_ELEMS: tl.constexpr, + KV_BLOCK_SIZE: tl.constexpr, + BLOCK: tl.constexpr, +): + """Snapshot one paged decode KV slot per batch row.""" + tile_id = tl.program_id(0) + batch_id = tl.program_id(1) + offs = tile_id * BLOCK + tl.arange(0, BLOCK) + valid_elem = offs < TOTAL_ELEMS + + head_id = offs // HEAD_DIM + dim_id = offs - head_id * HEAD_DIM + + kv_seqlen = tl.maximum(tl.load(kv_seqlens + batch_id), 1) + slot_id = kv_seqlen - 1 + block_idx = slot_id // KV_BLOCK_SIZE + page_offset = slot_id - block_idx * KV_BLOCK_SIZE + block_id = tl.load(block_offsets + batch_id * boff_stride_b + block_idx * boff_stride_n).to(tl.int64) + + k_ptrs = (k_cache + block_id * k_stride_n + page_offset * k_stride_b + head_id * k_stride_h + + dim_id * k_stride_d) + v_ptrs = (v_cache + block_id * v_stride_n + page_offset * v_stride_b + head_id * v_stride_h + + dim_id * v_stride_d) + sk_ptrs = saved_k + batch_id * sk_stride_b + head_id * sk_stride_h + dim_id * sk_stride_d + sv_ptrs = saved_v + batch_id * sv_stride_b + head_id * sv_stride_h + dim_id * sv_stride_d + tl.store(sk_ptrs, tl.load(k_ptrs, mask=valid_elem), mask=valid_elem) + tl.store(sv_ptrs, tl.load(v_ptrs, mask=valid_elem), mask=valid_elem) + + +@triton.jit +def _decode_kv_cache_restore_kernel( + k_cache, + v_cache, + saved_k, + saved_v, + block_offsets, + kv_seqlens, + restore_mask, + k_stride_n, + k_stride_b, + k_stride_h, + k_stride_d, + v_stride_n, + v_stride_b, + v_stride_h, + v_stride_d, + sk_stride_b, + sk_stride_h, + sk_stride_d, + sv_stride_b, + sv_stride_h, + sv_stride_d, + boff_stride_b, + boff_stride_n, + HEAD_DIM: tl.constexpr, + TOTAL_ELEMS: tl.constexpr, + KV_BLOCK_SIZE: tl.constexpr, + BLOCK: tl.constexpr, +): + """Restore one paged decode KV slot for masked batch rows.""" + tile_id = tl.program_id(0) + batch_id = tl.program_id(1) + do_restore = tl.load(restore_mask + batch_id) + if not do_restore: + return + + offs = tile_id * BLOCK + tl.arange(0, BLOCK) + valid_elem = offs < TOTAL_ELEMS + + head_id = offs // HEAD_DIM + dim_id = offs - head_id * HEAD_DIM + + kv_seqlen = tl.maximum(tl.load(kv_seqlens + batch_id), 1) + slot_id = kv_seqlen - 1 + block_idx = slot_id // KV_BLOCK_SIZE + page_offset = slot_id - block_idx * KV_BLOCK_SIZE + block_id = tl.load(block_offsets + batch_id * boff_stride_b + block_idx * boff_stride_n).to(tl.int64) + + k_ptrs = (k_cache + block_id * k_stride_n + page_offset * k_stride_b + head_id * k_stride_h + + dim_id * k_stride_d) + v_ptrs = (v_cache + block_id * v_stride_n + page_offset * v_stride_b + head_id * v_stride_h + + dim_id * v_stride_d) + sk_ptrs = saved_k + batch_id * sk_stride_b + head_id * sk_stride_h + dim_id * sk_stride_d + sv_ptrs = saved_v + batch_id * sv_stride_b + head_id * sv_stride_h + dim_id * sv_stride_d + tl.store(k_ptrs, tl.load(sk_ptrs, mask=valid_elem), mask=valid_elem) + tl.store(v_ptrs, tl.load(sv_ptrs, mask=valid_elem), mask=valid_elem) + + +@triton.jit +def _decode_concept_state_update_kernel( + last_raw_state_cache, + last_final_state_cache, + predicted_vectors, + raw_states, + state_ids, + update_mask, + raw_cache_stride_n, + raw_cache_stride_l, + raw_cache_stride_h, + final_cache_stride_n, + final_cache_stride_h, + pred_stride_b, + pred_stride_h, + raw_stride_b, + raw_stride_l, + raw_stride_h, + HIDDEN: tl.constexpr, + RAW_ELEMS: tl.constexpr, + BLOCK: tl.constexpr, +): + """Write final/raw concept state caches for valid boundary rows.""" + tile_id = tl.program_id(0) + batch_id = tl.program_id(1) + state_id = tl.load(state_ids + batch_id) + do_update = (state_id >= 0) & tl.load(update_mask + batch_id) + if not do_update: + return + + offs = tile_id * BLOCK + tl.arange(0, BLOCK) + hidden_id = offs % HIDDEN + + final_mask = offs < HIDDEN + pred_ptrs = predicted_vectors + batch_id * pred_stride_b + hidden_id * pred_stride_h + final_ptrs = last_final_state_cache + state_id * final_cache_stride_n + hidden_id * final_cache_stride_h + final_values = tl.load(pred_ptrs, mask=final_mask) + tl.store(final_ptrs, final_values, mask=final_mask) + + raw_mask = offs < RAW_ELEMS + layer_id = offs // HIDDEN + raw_ptrs = raw_states + batch_id * raw_stride_b + layer_id * raw_stride_l + hidden_id * raw_stride_h + raw_cache_ptrs = (last_raw_state_cache + state_id * raw_cache_stride_n + layer_id * raw_cache_stride_l + + hidden_id * raw_cache_stride_h) + raw_values = tl.load(raw_ptrs, mask=raw_mask) + tl.store(raw_cache_ptrs, raw_values, mask=raw_mask) + + def _flatten_decode_position_ids(position_ids: torch.Tensor, batch_size: int, device: torch.device) -> torch.Tensor: """Normalize decode position ids to one absolute position per batch row.""" if position_ids.dim() == 0: @@ -165,3 +326,137 @@ def decode_chunk_state_update( num_warps=8, ) return concept_inputs, next_rows, update_mask + + +def decode_kv_cache_snapshot( + k_cache: torch.Tensor, + v_cache: torch.Tensor, + block_offsets: torch.Tensor, + kv_seqlens: torch.Tensor, + block: int = 1024, +): + """Snapshot the current decode KV slot for each batch row. + + The slot is ``max(kv_seqlen, 1) - 1`` in the paged cache. This is used by + ConceptLM graph-safe decode to undo dummy all-row concept predictor writes + for non-boundary rows. + """ + assert k_cache.is_cuda and v_cache.is_cuda, 'ConceptLM KV snapshot requires CUDA caches.' + assert k_cache.dim() == 4 and v_cache.dim() == 4 + assert k_cache.shape[:3] == v_cache.shape[:3] + assert k_cache.shape[-1] == v_cache.shape[-1] + assert block_offsets.is_cuda and kv_seqlens.is_cuda + + batch_size = kv_seqlens.numel() + num_heads = k_cache.size(2) + head_dim = k_cache.size(3) + total_elems = num_heads * head_dim + saved_k = torch.empty((batch_size, num_heads, head_dim), dtype=k_cache.dtype, device=k_cache.device) + saved_v = torch.empty((batch_size, num_heads, head_dim), dtype=v_cache.dtype, device=v_cache.device) + grid = (triton.cdiv(total_elems, block), batch_size) + _decode_kv_cache_snapshot_kernel[grid]( + k_cache, + v_cache, + block_offsets, + kv_seqlens, + saved_k, + saved_v, + *k_cache.stride(), + *v_cache.stride(), + *block_offsets.stride(), + *saved_k.stride(), + *saved_v.stride(), + HEAD_DIM=head_dim, + TOTAL_ELEMS=total_elems, + KV_BLOCK_SIZE=k_cache.size(1), + BLOCK=block, + num_warps=8, + ) + return saved_k, saved_v + + +def decode_kv_cache_restore( + k_cache: torch.Tensor, + v_cache: torch.Tensor, + saved_k: torch.Tensor, + saved_v: torch.Tensor, + block_offsets: torch.Tensor, + kv_seqlens: torch.Tensor, + restore_mask: torch.Tensor, + block: int = 1024, +) -> None: + """Restore the current decode KV slot for masked batch rows.""" + assert k_cache.is_cuda and v_cache.is_cuda, 'ConceptLM KV restore requires CUDA caches.' + assert saved_k.is_cuda and saved_v.is_cuda + assert saved_k.shape == (kv_seqlens.numel(), k_cache.size(2), k_cache.size(3)) + assert saved_v.shape == (kv_seqlens.numel(), v_cache.size(2), v_cache.size(3)) + assert restore_mask.numel() == kv_seqlens.numel() + + batch_size = kv_seqlens.numel() + num_heads = k_cache.size(2) + head_dim = k_cache.size(3) + total_elems = num_heads * head_dim + restore_mask = restore_mask.to(device=k_cache.device, dtype=torch.bool) + grid = (triton.cdiv(total_elems, block), batch_size) + _decode_kv_cache_restore_kernel[grid]( + k_cache, + v_cache, + saved_k, + saved_v, + block_offsets, + kv_seqlens, + restore_mask, + *k_cache.stride(), + *v_cache.stride(), + *saved_k.stride(), + *saved_v.stride(), + *block_offsets.stride(), + HEAD_DIM=head_dim, + TOTAL_ELEMS=total_elems, + KV_BLOCK_SIZE=k_cache.size(1), + BLOCK=block, + num_warps=8, + ) + + +def decode_concept_state_update( + last_raw_state_cache: torch.Tensor, + last_final_state_cache: torch.Tensor, + predicted_vectors: torch.Tensor, + raw_states: torch.Tensor, + state_ids: torch.Tensor, + update_mask: torch.Tensor, + block: int = 1024, +) -> None: + """Write final/raw concept states for valid boundary rows.""" + assert last_raw_state_cache.is_cuda, 'ConceptLM concept-state update requires CUDA caches.' + assert last_final_state_cache.is_cuda + assert predicted_vectors.is_cuda and raw_states.is_cuda + assert raw_states.dim() == 3 + assert last_raw_state_cache.shape[1:] == raw_states.shape[1:] + assert last_final_state_cache.size(1) == predicted_vectors.size(1) + assert predicted_vectors.size(0) == raw_states.size(0) == state_ids.numel() == update_mask.numel() + + batch_size = predicted_vectors.size(0) + hidden = predicted_vectors.size(1) + raw_elems = raw_states.size(1) * raw_states.size(2) + max_elems = max(hidden, raw_elems) + state_ids = state_ids.to(device=predicted_vectors.device, dtype=torch.long) + update_mask = update_mask.to(device=predicted_vectors.device, dtype=torch.bool) + grid = (triton.cdiv(max_elems, block), batch_size) + _decode_concept_state_update_kernel[grid]( + last_raw_state_cache, + last_final_state_cache, + predicted_vectors, + raw_states, + state_ids, + update_mask, + *last_raw_state_cache.stride(), + *last_final_state_cache.stride(), + *predicted_vectors.stride(), + *raw_states.stride(), + HIDDEN=hidden, + RAW_ELEMS=raw_elems, + BLOCK=block, + num_warps=8, + ) diff --git a/lmdeploy/pytorch/models/intern_ncp.py b/lmdeploy/pytorch/models/intern_ncp.py index f8121a515d..6a4f30a3be 100644 --- a/lmdeploy/pytorch/models/intern_ncp.py +++ b/lmdeploy/pytorch/models/intern_ncp.py @@ -1,29 +1,4 @@ # Copyright (c) OpenMMLab. All rights reserved. -"""lmdeploy adapter for ConceptLM V2.2-VQ. - -Reference implementation: -``concept_olmo_stage_2_V1/modeling_conceptlm_v22_vq.py`` - -Modules are added incrementally. Current state: - - token embedding + output projection (lm_head) - - ``_OlmoBlock`` (encoder/decoder/concept_predictor backbone): attention, - mlp, rmsnorm, rope. Wired to lmdeploy primitives so it is TP-correct and - ready to plug into the engine's paged attention path. - - ``_Quantizer``: stacked VQ codebook parameter, replicated across TP. - - ``_SelfDD``: replicated per-token depth mixer for encoder hidden history. - - ``_ResidualRoute``: replicated residual source mixer for decoder routes. - - ``_TwoRouteAdd``: decoder depth mixing plus final-concept residual route. - - ``_ConceptPredictor``: concept block container and prediction heads. - - top-level encoder/decoder containers, fusion norms, route norms, and - checkpoint loading for the implemented module tree. - - packed non-decode prefill path through encoder -> concept - predictor/quantizer -> fusion -> decoder, with per-request chunk-stream - attention metadata derived explicitly from token-stream metadata. - -Decode still needs the graph-safe compressed concept-stream runtime contract. -The module structure mirrors the reference for readability. -""" - from collections.abc import Iterable from dataclasses import dataclass, replace from typing import Any @@ -34,10 +9,21 @@ from transformers.configuration_utils import PretrainedConfig from lmdeploy.pytorch.model_inputs import StepContext, StepContextManager -from lmdeploy.pytorch.nn import (ApplyRotaryEmb, Attention, ConceptLMRuntimeOps, RMSNorm, SiluAndMul, - build_rotary_embedding) -from lmdeploy.pytorch.nn.linear import (build_down_linear, build_gateup_linear, - build_merged_colwise_linear, build_o_proj, build_qkv_proj) +from lmdeploy.pytorch.nn import ( + ApplyRotaryEmb, + Attention, + ConceptLMRuntimeOps, + RMSNorm, + SiluAndMul, + build_rotary_embedding, +) +from lmdeploy.pytorch.nn.linear import ( + build_down_linear, + build_gateup_linear, + build_merged_colwise_linear, + build_o_proj, + build_qkv_proj, +) from lmdeploy.pytorch.weight_loader.model_weight_loader import load_weight from .patch import add_prefix @@ -263,7 +249,8 @@ def _make_olmo_rotary_embedding(config: PretrainedConfig, def _repack_olmo_qkv_weight(loaded_weight: torch.Tensor, num_heads: int, head_dim: int): - """Convert native OLMo per-head [Q,K,V] QKV packing to LMDeploy [Q][K][V].""" + """Convert native OLMo per-head [Q,K,V] QKV packing to LMDeploy + [Q][K][V].""" leading_shape = loaded_weight.shape[1:] loaded_weight = loaded_weight.reshape(num_heads, 3, head_dim, *leading_shape) query = loaded_weight[:, 0].flatten(0, 1) @@ -273,7 +260,8 @@ def _repack_olmo_qkv_weight(loaded_weight: torch.Tensor, num_heads: int, head_di def _load_stacked_codebook_weight(param: torch.nn.Parameter, loaded_weight: torch.Tensor, codebook_idx: int): - """Load one native ``codebook.N`` checkpoint tensor into a stacked codebook.""" + """Load one native ``codebook.N`` checkpoint tensor into a stacked + codebook.""" assert 0 <= codebook_idx < param.size(0), f'Invalid codebook index: {codebook_idx}' target = param.data[codebook_idx] assert target.size() == loaded_weight.size(), ( @@ -487,7 +475,8 @@ def __init__(self, ) def transformed_codebook(self): - """Return codebook as ``[num_codebooks, codebook_size, codebook_dim]``.""" + """Return codebook as ``[num_codebooks, codebook_size, + codebook_dim]``.""" return self.codebook def forward(self, concept_logits: torch.Tensor) -> torch.Tensor: @@ -590,7 +579,8 @@ def __init__(self, ]) def make_history_buffer(self, hidden_states: torch.Tensor) -> torch.Tensor: - """Allocate layer-major history buffer ``[num_layers + 1, *hidden_shape]``.""" + """Allocate layer-major history buffer ``[num_layers + 1, + *hidden_shape]``.""" return hidden_states.new_empty((self.num_layers + 1, *hidden_states.shape)) @staticmethod @@ -667,7 +657,8 @@ def _source_tensor(self, source_states: _SourceStates, source_dim: int, expected_leading_shape: tuple[int, ...] | None = None): - """Return source states in shape ``[..., active_sources, hidden_size]``.""" + """Return source states in shape ``[..., active_sources, + hidden_size]``.""" if source_states is None: return None if not isinstance(source_states, torch.Tensor): @@ -710,7 +701,8 @@ def _add_update(self, return target_hidden + update.to(target_hidden.dtype) def make_source_buffer(self, hidden_states: torch.Tensor) -> torch.Tensor: - """Allocate source-major buffer ``[num_source_states, *hidden_shape]``.""" + """Allocate source-major buffer ``[num_source_states, + *hidden_shape]``.""" return hidden_states.new_empty((self.num_source_states, *hidden_states.shape)) @staticmethod @@ -875,7 +867,8 @@ def forward_repeated_chunks_active_from_buffer(self, shift_feature: bool, active_sources: int, residual_scale: torch.Tensor | None = None): - """Flexible/debug repeated-chunk buffer path with reduced active source count.""" + """Flexible/debug repeated-chunk buffer path with reduced active source + count.""" return self.forward_repeated_chunks( target_hidden, self.source_view(source_buffer, active_sources), @@ -889,9 +882,8 @@ def forward_repeated_chunks_active_from_buffer(self, class ConceptLMV22VQConceptRoute(nn.Module): """Rewrite of ``_ConceptRoute``. - Applies LayerNorm to the final concept state, scales it elementwise with a - learned diagonal, optionally applies a route scale, then adds the update to - decoder hidden states. This is replicated and token-local. + Applies LayerNorm to the final concept state, scales it elementwise with a learned diagonal, optionally applies a + route scale, then adds the update to decoder hidden states. This is replicated and token-local. """ def __init__(self, @@ -944,7 +936,8 @@ def __init__(self, ]) def make_history_buffer(self, hidden_states: torch.Tensor) -> torch.Tensor: - """Allocate layer-major decoder history buffer ``[num_layers + 1, *hidden_shape]``.""" + """Allocate layer-major decoder history buffer ``[num_layers + 1, + *hidden_shape]``.""" return hidden_states.new_empty((self.num_layers + 1, *hidden_states.shape)) @staticmethod @@ -954,7 +947,8 @@ def write_history(history_buffer: torch.Tensor, slot_idx: int, hidden_states: to @staticmethod def history_view(history_buffer: torch.Tensor, layer_idx: int): - """Return layer-major decoder history needed by ``layer_idx`` without copying.""" + """Return layer-major decoder history needed by ``layer_idx`` without + copying.""" return ConceptLMV22VQSelfDD.history_view(history_buffer, layer_idx) def forward_from_buffer(self, @@ -1028,10 +1022,8 @@ def forward(self, hidden_states: torch.Tensor): class ConceptLMV22VQConceptPredictor(nn.Module): """Rewrite of ``_ConceptPredictor``. - The predictor owns the high-level concept OLMo block, per-codebook - prediction heads, concept self-DD, encoder-read routes, and the shared - source LayerNorm used before encoder states are routed into the concept - stream. + The predictor owns the high-level concept OLMo block, per-codebook prediction heads, concept self-DD, encoder-read + routes, and the shared source LayerNorm used before encoder states are routed into the concept stream. """ def __init__(self, @@ -1085,7 +1077,8 @@ def set_attention_window(self, window_size, skip_frequency): self._window_skip_frequency = skip_frequency def normalize_encoder_concept_states(self, encoder_concept_states: torch.Tensor): - """Apply the shared source norm used before concept-read-encoder routes.""" + """Apply the shared source norm used before concept-read-encoder + routes.""" return self.concept_read_encoder_shared_source_norm(encoder_concept_states) def predict_logits(self, hidden_states: torch.Tensor): @@ -1490,7 +1483,8 @@ def forward(self, return hidden_states, layer_states def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]], prefix: str = ''): - """Load native ConceptLM OLMo block weights into the LMDeploy rewrite.""" + """Load native ConceptLM OLMo block weights into the LMDeploy + rewrite.""" if prefix and not prefix.endswith('.'): prefix = f'{prefix}.' @@ -1722,9 +1716,8 @@ def _decode_chunk_state_update(self, concept_caches: ConceptCaches) -> ConceptChunkStateUpdateResult: """Update decode chunk-source state and return fixed-shape rows. - CUDA uses the Triton writer. CPU uses the reference writer for tests. - The returned rows deliberately avoid dynamic concept-row compaction, - matching the CUDA graph route in the design doc. + CUDA uses the Triton writer. CPU uses the reference writer for tests. The returned rows deliberately avoid + dynamic concept-row compaction, matching the CUDA graph route in the design doc. """ chunk_source_state = concept_caches.chunk_source_state if chunk_source_state is None: @@ -1760,18 +1753,18 @@ def support_cuda_graph( inputs_embeds: torch.Tensor = None, **kwargs, ): - """Disable CUDA graph until ConceptLM concept-boundary updates are graph-safe. + """Enable CUDA graph for fixed-shape decode only. - ``states_shapes`` makes the engine allocate graph-padded state ids. The - eager decode path updates chunk-source state through a backend op, but - still compacts boundary concept rows dynamically before concept - predictor attention. Capturing that would bake in a batch-specific - concept update shape. + Prefill still builds packed concept metadata with runtime compact sizes + and must stay eager. Decode boundary updates use fixed batch-shape + backend ops, so the CUDA graph runner can pad ``state_ids`` through the + common SSM path and replay the model with static tensor shapes. """ - return False + return bool(getattr(attn_metadata, 'is_decoding', False)) def _route_gate(self, layer_idx: int) -> torch.Tensor: - """Return decoder route gate ``[decoder_dd_scale, concept_route_scale]``.""" + """Return decoder route gate ``[decoder_dd_scale, + concept_route_scale]``.""" return self.final_read_concept_gate_logits[int(layer_idx)].float().softmax(dim=-1) def _normalize_prefill_inputs(self, @@ -1811,7 +1804,8 @@ def _normalize_prefill_inputs(self, return hidden_states, position_ids[0].to(device=hidden_states.device) def _normalize_decode_inputs(self, hidden_states: torch.Tensor, position_ids: torch.Tensor): - """Normalize LMDeploy decode inputs to one flat row per active request.""" + """Normalize LMDeploy decode inputs to one flat row per active + request.""" if hidden_states.dim() != 3 or hidden_states.size(0) != 1: raise NotImplementedError( f'ConceptLM decode expects fixed engine layout [1, batch, hidden], ' @@ -1851,7 +1845,8 @@ def _select_decode_state_rows(state_cache: torch.Tensor, return torch.where(valid_mask, rows, torch.zeros_like(rows)) def _decode_concept_read_mask(self, decode_metadata: ConceptDecodeMetadata) -> torch.Tensor: - """Return rows whose current decode token should read a cached concept.""" + """Return rows whose current decode token should read a cached + concept.""" repeat_slots = self._repeat_slot_ids( decode_metadata.position_ids, int(self.config.concept_chunk_size), @@ -1862,7 +1857,8 @@ def _decode_concept_read_mask(self, decode_metadata: ConceptDecodeMetadata) -> t def _build_decode_chunk_source_states(self, hidden_states: torch.Tensor, encoder_raw_states: list[torch.Tensor]) -> torch.Tensor: - """Build current per-row states accumulated until the next concept boundary. + """Build current per-row states accumulated until the next concept + boundary. Row 0 is the final encoder hidden, used as concept-predictor input when a boundary is reached. Remaining rows mirror the prefill @@ -1874,26 +1870,25 @@ def _build_decode_chunk_source_states(self, @staticmethod def _decode_concept_position_ids(position_ids: torch.Tensor, chunk_size: int) -> torch.Tensor: - """Return reference RoPE positions for concept rows emitted at decode boundaries.""" + """Return reference RoPE positions for concept rows emitted at decode + boundaries.""" return (position_ids - int(chunk_size) + 1).clamp(min=0) - def _build_concept_decode_metadata_eager(self, - token_attn_metadata: Any, - decode_metadata: ConceptDecodeMetadata, - boundary_indices: torch.Tensor): - """Build dynamic concept-stream decode metadata for boundary rows. - - This is the eager-only bridge: it compacts rows whose token completed a - concept chunk and makes concept KV positions advance on the compressed - concept timeline. CUDA graph support should replace this with a - fixed-shape backend metadata object/op instead of calling ``nonzero`` - and constructing variable-batch attention metadata here. + def _build_concept_decode_metadata_static(self, + token_attn_metadata: Any, + decode_metadata: ConceptDecodeMetadata): + """Build fixed-shape concept-stream decode metadata. + + Concept predictor runs with the same batch shape as token decode. + Boundary rows append a real concept KV entry. Non-boundary/padded rows + execute dummy concept attention with safe ``kv_seqlens >= 1``; their KV + writes are restored afterward by ``ConceptLMRuntimeOps``. """ if token_attn_metadata is None or getattr(token_attn_metadata, 'block_offsets', None) is None: raise RuntimeError('ConceptLM decode concept update requires token attention metadata.') device = decode_metadata.position_ids.device - num_boundary_rows = boundary_indices.numel() + batch_size = decode_metadata.position_ids.numel() q_seqlens = getattr(token_attn_metadata, 'q_seqlens', None) q_start_loc = getattr(token_attn_metadata, 'q_start_loc', None) kv_seqlens = getattr(token_attn_metadata, 'kv_seqlens', None) @@ -1903,18 +1898,18 @@ def _build_concept_decode_metadata_eager(self, q_start_dtype = q_start_loc.dtype kv_dtype = kv_seqlens.dtype - concept_q_seqlens = torch.ones((num_boundary_rows, ), dtype=q_dtype, device=device) - concept_q_start_loc = torch.arange(num_boundary_rows, dtype=q_start_dtype, device=device) + concept_q_seqlens = torch.ones((batch_size, ), dtype=q_dtype, device=device) + concept_q_start_loc = torch.arange(batch_size, dtype=q_start_dtype, device=device) concept_cu_seqlens = F.pad(torch.cumsum(concept_q_seqlens, dim=0, dtype=torch.int32), (1, 0)) concept_kv_seqlens = torch.div( - decode_metadata.position_ids.index_select(0, boundary_indices) + 1, + decode_metadata.position_ids + 1, int(self.config.concept_chunk_size), rounding_mode='floor', - ).to(dtype=kv_dtype) + ).clamp(min=1).to(dtype=kv_dtype) updates = dict( is_decoding=True, - block_offsets=token_attn_metadata.block_offsets.index_select(0, boundary_indices), + block_offsets=token_attn_metadata.block_offsets, q_start_loc=concept_q_start_loc, q_seqlens=concept_q_seqlens, kv_seqlens=concept_kv_seqlens, @@ -1924,7 +1919,7 @@ def _build_concept_decode_metadata_eager(self, if hasattr(token_attn_metadata, 'kv_start_loc'): updates['kv_start_loc'] = concept_kv_seqlens - concept_q_seqlens.to(dtype=concept_kv_seqlens.dtype) if hasattr(token_attn_metadata, 'kv_flatten_size'): - updates['kv_flatten_size'] = num_boundary_rows + updates['kv_flatten_size'] = batch_size if hasattr(token_attn_metadata, 'max_q_seqlen'): updates['max_q_seqlen'] = 1 if hasattr(token_attn_metadata, 'max_kv_seqlen'): @@ -1943,56 +1938,94 @@ def _build_concept_decode_metadata_eager(self, @staticmethod def _stack_concept_raw_states(concept_raw_states: list[torch.Tensor]) -> torch.Tensor: - """Stack raw concept-layer states to ``[rows, concept_layers, hidden]``.""" + """Stack raw concept-layer states to ``[rows, concept_layers, + hidden]``.""" return torch.stack(tuple(concept_raw_states), dim=1) - def _write_decode_concept_states_eager_(self, - concept_caches: ConceptCaches, - decode_metadata: ConceptDecodeMetadata, - boundary_indices: torch.Tensor, - predicted_vectors: torch.Tensor, - concept_raw_states: list[torch.Tensor]): + def _snapshot_decode_concept_kv(self, + concept_caches: ConceptCaches, + concept_attn_metadata: Any): + """Snapshot concept KV slots that dummy non-boundary rows may + overwrite.""" + if concept_caches.concept_past_key_values is None: + raise RuntimeError('ConceptLM decode concept update requires concept KV caches.') + return [ + self.concept_ops.decode_kv_cache_snapshot( + k_cache, + v_cache, + concept_attn_metadata.block_offsets, + concept_attn_metadata.kv_seqlens, + ) + for k_cache, v_cache in concept_caches.concept_past_key_values + ] + + def _restore_decode_concept_kv_( + self, + concept_caches: ConceptCaches, + concept_attn_metadata: Any, + saved_kv, + restore_mask: torch.Tensor, + ): + """Restore concept KV slots for non-boundary and padded rows.""" + if concept_caches.concept_past_key_values is None: + raise RuntimeError('ConceptLM decode concept update requires concept KV caches.') + for (k_cache, v_cache), (saved_k, saved_v) in zip(concept_caches.concept_past_key_values, saved_kv): + self.concept_ops.decode_kv_cache_restore( + k_cache, + v_cache, + saved_k, + saved_v, + concept_attn_metadata.block_offsets, + concept_attn_metadata.kv_seqlens, + restore_mask, + ) + + def _write_decode_concept_states_static_(self, + concept_caches: ConceptCaches, + decode_metadata: ConceptDecodeMetadata, + update_mask: torch.Tensor, + predicted_vectors: torch.Tensor, + concept_raw_states: list[torch.Tensor]): """Write newly emitted concept states to persistent decode caches.""" last_final_state = concept_caches.last_final_state last_raw_states = concept_caches.last_raw_states if last_final_state is None or last_raw_states is None: raise RuntimeError('ConceptLM decode concept update requires last concept state caches.') - state_ids = decode_metadata.safe_state_ids.index_select(0, boundary_indices) - last_final_state.index_copy_(0, state_ids, predicted_vectors.to(dtype=last_final_state.dtype)) raw_rows = self._stack_concept_raw_states(concept_raw_states) - last_raw_states.index_copy_(0, state_ids, raw_rows.to(dtype=last_raw_states.dtype)) + self.concept_ops.decode_concept_state_update( + last_raw_states, + last_final_state, + predicted_vectors, + raw_rows, + decode_metadata.state_ids, + update_mask, + ) - def _update_decode_concept_states_eager_(self, - chunk_update: ConceptChunkStateUpdateResult, - decode_metadata: ConceptDecodeMetadata, - concept_metadata: ConceptMetadata, - concept_caches: ConceptCaches): - """Emit and cache concept states for decode rows that complete a chunk. - - TODO: replace this dynamic eager bridge with a graph-safe backend op and - fixed-shape concept attention metadata. The model-level flow should stay - the same: chunk accumulator -> concept predictor on boundary rows -> - cached final/raw concept states -> decoder routes. + def _update_decode_concept_states_static_(self, + chunk_update: ConceptChunkStateUpdateResult, + decode_metadata: ConceptDecodeMetadata, + concept_metadata: ConceptMetadata, + concept_caches: ConceptCaches): + """Emit/cache concept states with fixed batch shape. + + The predictor runs for every decode row so CUDA graph capture sees a stable launch sequence. Non-boundary rows + are dummy work: their concept KV writes are restored and their final/raw state writes are masked. """ - boundary_indices = torch.nonzero(chunk_update.concept_update_mask, as_tuple=False).flatten() - if boundary_indices.numel() == 0: - return if concept_caches.concept_past_key_values is None: raise RuntimeError('ConceptLM decode concept update requires concept KV caches.') - boundary_concept_inputs = chunk_update.concept_input_states.index_select(0, boundary_indices) - concept_hidden = self.concept_vq_input_norm(boundary_concept_inputs[:, 0]) + concept_hidden = self.concept_vq_input_norm(chunk_update.concept_input_states[:, 0]) encoder_concept_states = self.concept_predictor.normalize_encoder_concept_states( - boundary_concept_inputs[:, 1:]) + chunk_update.concept_input_states[:, 1:]) concept_position_ids = self._decode_concept_position_ids( decode_metadata.position_ids, concept_metadata.chunk_size, - ).index_select(0, boundary_indices) - concept_attn_metadata = self._build_concept_decode_metadata_eager( + ) + concept_attn_metadata = self._build_concept_decode_metadata_static( concept_metadata.attn_metadata, decode_metadata, - boundary_indices, ) + saved_kv = self._snapshot_decode_concept_kv(concept_caches, concept_attn_metadata) concept_logits, concept_raw_states = self.concept_predictor( concept_hidden, encoder_concept_states, @@ -2000,11 +2033,17 @@ def _update_decode_concept_states_eager_(self, past_key_values=concept_caches.concept_past_key_values, attn_metadata=concept_attn_metadata, ) + self._restore_decode_concept_kv_( + concept_caches, + concept_attn_metadata, + saved_kv, + ~chunk_update.concept_update_mask, + ) predicted_vectors = self.concept_quantizer(concept_logits) - self._write_decode_concept_states_eager_( + self._write_decode_concept_states_static_( concept_caches, decode_metadata, - boundary_indices, + chunk_update.concept_update_mask, predicted_vectors, concept_raw_states, ) @@ -2112,7 +2151,8 @@ def _concept_counts_from_q_seqlens(q_seqlens: torch.Tensor, chunk_size: int) -> @staticmethod def _repeat_slot_ids(token_pos: torch.Tensor, chunk_size: int, shift_feature: bool) -> torch.Tensor: - """Return local concept slot read by each token after shift semantics.""" + """Return local concept slot read by each token after shift + semantics.""" if shift_feature: return torch.div(token_pos + 1, chunk_size, rounding_mode='floor') - 1 return torch.div(token_pos, chunk_size, rounding_mode='floor') - 1 @@ -2121,7 +2161,8 @@ def _repeat_slot_ids(token_pos: torch.Tensor, chunk_size: int, shift_feature: bo def _get_max_concepts_per_request(token_attn_metadata: Any, concept_q_seqlens: torch.Tensor, chunk_size: int) -> int: - """Return per-request concept attention bound without hidden context access.""" + """Return per-request concept attention bound without hidden context + access.""" max_q_seqlen = getattr(token_attn_metadata, 'max_q_seqlen', None) if max_q_seqlen is not None: return ConceptLMV22VQForCausalLM._concept_count_from_seq_len(int(max_q_seqlen), chunk_size) @@ -2290,7 +2331,8 @@ def _repeat_shift_source_states_packed(self, def _build_encoder_concept_states_packed(self, encoder_raw_states: list[torch.Tensor], prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: - """Build packed chunk-level encoder states used by the concept predictor.""" + """Build packed chunk-level encoder states used by the concept + predictor.""" chunks = [self._merge_chunks_packed(state, prefill_metadata) for state in encoder_raw_states[:-1]] assert len(chunks) > 0, 'ConceptLM concept-read-encoder route requires at least one encoder source state.' states = torch.stack(chunks, dim=-2) @@ -2390,12 +2432,11 @@ def _forward_decode(self, position_ids: torch.Tensor, concept_metadata: ConceptMetadata, concept_caches: ConceptCaches): - """Eager ConceptLM decode path. + """ConceptLM decode path. - This path is semantically structured for serving but intentionally keeps - CUDA graph disabled. Boundary concept updates currently use dynamic row - compaction in ``_update_decode_concept_states_eager_``; the stable - graph route needs a fixed-shape concept metadata/backend op. + This path is semantically structured for serving. Boundary concept + updates run with fixed batch shape so it is eligible for CUDA graph + replay through ``support_cuda_graph``. """ if concept_caches.encoder_past_key_values is None or concept_caches.decoder_past_key_values is None: raise RuntimeError('ConceptLM decode requires encoder and decoder KV caches.') @@ -2432,7 +2473,7 @@ def _forward_decode(self, decode_concept_metadata, concept_caches, ) - self._update_decode_concept_states_eager_( + self._update_decode_concept_states_static_( chunk_update, decode_metadata, decode_concept_metadata, @@ -2471,10 +2512,9 @@ def _encode(self, attn_metadata: Any = None): """Encoder stack plus encoder self-DD. - This helper mirrors the reference flow but consumes LMDeploy attention - inputs. It is wired for the future full forward path; continuous - batching still needs caller-side concept metadata before the full model - can use it safely. + This helper mirrors the reference flow but consumes LMDeploy attention inputs. It is wired for the future full + forward path; continuous batching still needs caller-side concept metadata before the full model can use it + safely. """ raw_states = [] history_buffer = self.dd_encoder_self_dd.make_history_buffer(hidden_states) @@ -2601,7 +2641,8 @@ def prepare_inputs_for_generation(self, ) def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]): - """Load native ConceptLM checkpoint weights into implemented modules.""" + """Load native ConceptLM checkpoint weights into implemented + modules.""" # (checkpoint_name, target_name) weight_map = { 'embedding.word_embeddings.weight': 'embedding.word_embeddings.weight', diff --git a/lmdeploy/pytorch/nn/conceptlm.py b/lmdeploy/pytorch/nn/conceptlm.py index 1aeb2a505f..22f12bcf31 100644 --- a/lmdeploy/pytorch/nn/conceptlm.py +++ b/lmdeploy/pytorch/nn/conceptlm.py @@ -7,8 +7,8 @@ class ConceptLMRuntimeOps(nn.Module): """ConceptLM model-specific runtime operation wrapper. - The model calls this nn module only. Backend implementations own dispatch, - and CUDA implementations own direct Triton kernel launchers. + The model calls this nn module only. Backend implementations own dispatch, and CUDA implementations own direct + Triton kernel launchers. """ def __init__(self): @@ -26,7 +26,8 @@ def decode_chunk_state_update( chunk_size: int, merge_method: str, ) -> tuple[Tensor, Tensor, Tensor]: - """Update state cache and return concept inputs, next rows, and mask.""" + """Update state cache and return concept inputs, next rows, and + mask.""" return self.impl.decode_chunk_state_update( chunk_source_state_cache, current_source_states, @@ -36,6 +37,61 @@ def decode_chunk_state_update( merge_method, ) + def decode_kv_cache_snapshot( + self, + k_cache: Tensor, + v_cache: Tensor, + block_offsets: Tensor, + kv_seqlens: Tensor, + ) -> tuple[Tensor, Tensor]: + """Snapshot one decode KV slot per batch row.""" + return self.impl.decode_kv_cache_snapshot( + k_cache, + v_cache, + block_offsets, + kv_seqlens, + ) + + def decode_kv_cache_restore( + self, + k_cache: Tensor, + v_cache: Tensor, + saved_k: Tensor, + saved_v: Tensor, + block_offsets: Tensor, + kv_seqlens: Tensor, + restore_mask: Tensor, + ) -> None: + """Restore one decode KV slot for masked batch rows.""" + return self.impl.decode_kv_cache_restore( + k_cache, + v_cache, + saved_k, + saved_v, + block_offsets, + kv_seqlens, + restore_mask, + ) + + def decode_concept_state_update( + self, + last_raw_state_cache: Tensor, + last_final_state_cache: Tensor, + predicted_vectors: Tensor, + raw_states: Tensor, + state_ids: Tensor, + update_mask: Tensor, + ) -> None: + """Write final/raw concept states for masked decode rows.""" + return self.impl.decode_concept_state_update( + last_raw_state_cache, + last_final_state_cache, + predicted_vectors, + raw_states, + state_ids, + update_mask, + ) + def forward( self, chunk_source_state_cache: Tensor, From 15d53d43e5f9b7e5bb07c19fb597102790cc063a Mon Sep 17 00:00:00 2001 From: grimoire Date: Sun, 26 Jul 2026 15:53:34 +0800 Subject: [PATCH 06/16] Refactor ConceptLM intern ncp implementation --- lmdeploy/pytorch/backends/conceptlm.py | 5 +- lmdeploy/pytorch/backends/cuda/conceptlm.py | 4 +- .../pytorch/backends/default/conceptlm.py | 6 +- lmdeploy/pytorch/configurations/conceptlm.py | 37 +- lmdeploy/pytorch/kernels/cuda/conceptlm.py | 19 +- lmdeploy/pytorch/models/intern_ncp.py | 2697 ----------------- .../pytorch/models/intern_ncp/__init__.py | 6 + .../pytorch/models/intern_ncp/metadata.py | 278 ++ .../pytorch/models/intern_ncp/modeling.py | 1369 +++++++++ lmdeploy/pytorch/models/intern_ncp/modules.py | 984 ++++++ lmdeploy/pytorch/models/intern_ncp/weight.py | 23 + lmdeploy/pytorch/nn/conceptlm.py | 24 +- 12 files changed, 2693 insertions(+), 2759 deletions(-) delete mode 100644 lmdeploy/pytorch/models/intern_ncp.py create mode 100644 lmdeploy/pytorch/models/intern_ncp/__init__.py create mode 100644 lmdeploy/pytorch/models/intern_ncp/metadata.py create mode 100644 lmdeploy/pytorch/models/intern_ncp/modeling.py create mode 100644 lmdeploy/pytorch/models/intern_ncp/modules.py create mode 100644 lmdeploy/pytorch/models/intern_ncp/weight.py diff --git a/lmdeploy/pytorch/backends/conceptlm.py b/lmdeploy/pytorch/backends/conceptlm.py index 5aec30e84d..f6c279c724 100644 --- a/lmdeploy/pytorch/backends/conceptlm.py +++ b/lmdeploy/pytorch/backends/conceptlm.py @@ -20,9 +20,8 @@ def decode_chunk_state_update( position_ids: Tensor, chunk_size: int, merge_method: str, - ) -> tuple[Tensor, Tensor, Tensor]: - """Update state cache and return concept inputs, next rows, and - mask.""" + ) -> tuple[Tensor, Tensor]: + """Update state cache and return concept inputs plus update mask.""" raise NotImplementedError('Not implemented.') @abstractmethod diff --git a/lmdeploy/pytorch/backends/cuda/conceptlm.py b/lmdeploy/pytorch/backends/cuda/conceptlm.py index 61672957f2..6957ba910e 100644 --- a/lmdeploy/pytorch/backends/cuda/conceptlm.py +++ b/lmdeploy/pytorch/backends/cuda/conceptlm.py @@ -22,8 +22,8 @@ def decode_chunk_state_update( position_ids: Tensor, chunk_size: int, merge_method: str, - ) -> tuple[Tensor, Tensor, Tensor]: - """Update state cache and return fixed-shape decode rows.""" + ) -> tuple[Tensor, Tensor]: + """Update state cache and return fixed-shape concept inputs.""" return decode_chunk_state_update( chunk_source_state_cache, current_source_states, diff --git a/lmdeploy/pytorch/backends/default/conceptlm.py b/lmdeploy/pytorch/backends/default/conceptlm.py index ada11b9866..c72906d5d0 100644 --- a/lmdeploy/pytorch/backends/default/conceptlm.py +++ b/lmdeploy/pytorch/backends/default/conceptlm.py @@ -43,8 +43,8 @@ def decode_chunk_state_update( position_ids: Tensor, chunk_size: int, merge_method: str, - ) -> tuple[Tensor, Tensor, Tensor]: - """Update state cache and return fixed-shape decode rows.""" + ) -> tuple[Tensor, Tensor]: + """Update state cache and return fixed-shape concept inputs.""" assert current_source_states.dim() == 3, ( f'current_source_states must be [batch, num_sources, hidden], got {tuple(current_source_states.shape)}.') assert chunk_source_state_cache.dim() == 3, ( @@ -89,7 +89,7 @@ def decode_chunk_state_update( state_id = int(state_ids[batch_idx]) if state_id >= 0: chunk_source_state_cache[state_id].copy_(next_rows[batch_idx]) - return concept_input_states, next_rows, update_mask + return concept_input_states, update_mask def decode_kv_cache_snapshot( self, diff --git a/lmdeploy/pytorch/configurations/conceptlm.py b/lmdeploy/pytorch/configurations/conceptlm.py index c66b269ac9..5ce5b75ee3 100644 --- a/lmdeploy/pytorch/configurations/conceptlm.py +++ b/lmdeploy/pytorch/configurations/conceptlm.py @@ -1,6 +1,7 @@ # Copyright (c) OpenMMLab. All rights reserved. import torch +from lmdeploy.pytorch.config import StateCacheSpec from lmdeploy.utils import get_logger from .builder import AutoModelConfigBuilder @@ -9,12 +10,10 @@ logger = get_logger('lmdeploy') CONCEPT_STATE_CHUNK_SOURCE = 0 -CONCEPT_STATE_LAST_RAW = 1 -CONCEPT_STATE_LAST_FINAL = 2 +CONCEPT_STATE_LAST = 1 CONCEPT_STATE_NAMES = ( 'concept_chunk_source_state', - 'concept_last_raw_states', - 'concept_last_final_state', + 'concept_last_state', ) @@ -34,9 +33,8 @@ def _get_concept_state_dtype(hf_config): class ConceptLMModelConfigBuilder(AutoModelConfigBuilder): """Config builder for ConceptLM V2.2-VQ. - The upstream checkpoint config does not declare bos/eos/pad token ids, - so derive them from the tokenizer when missing before handing the config - to the default builder (which asserts they exist). + The upstream checkpoint config does not declare bos/eos/pad token ids, so derive them from the tokenizer when + missing before handing the config to the default builder (which asserts they exist). """ @classmethod @@ -77,19 +75,24 @@ def build(cls, hf_config, model_path: str = None, **kwargs): # used as concept-predictor input; following rows are encoder raw states # consumed by concept-read-encoder residual routes. concept_chunk_state_sources = 1 + concept_encoder_read_sources - model_config.states_shapes = [ - ((concept_chunk_state_sources, hidden_size), state_dtype), - ((concept_layers, hidden_size), state_dtype), - ((hidden_size, ), state_dtype), + # Last emitted concept snapshot. Row 0 is the final concept vector, + # rows 1: are raw concept-layer states. Keep this packed so decode can + # gather the visible last-concept state once per batch row. + concept_last_state_sources = 1 + concept_layers + state_specs = [ + StateCacheSpec(CONCEPT_STATE_NAMES[CONCEPT_STATE_CHUNK_SOURCE], + (concept_chunk_state_sources, hidden_size), state_dtype), + StateCacheSpec(CONCEPT_STATE_NAMES[CONCEPT_STATE_LAST], (concept_last_state_sources, hidden_size), + state_dtype), ] - # The current branch only supports anonymous states_shapes. Keep stable - # indices on the HF config so model code has one semantic source of - # truth; if DSV4 StateCacheSpec lands here later, these become the - # names of ConceptLM's state-cache specs. + model_config.state_cache_specs = state_specs + # Backward-compat bridge used by scheduler/state-cache sizing. The + # actual runtime access should use state_cache_specs/named_state_caches + # like DSV4, not anonymous order-dependent indices. + model_config.states_shapes = [(tuple(spec.shape), spec.dtype) for spec in state_specs] model_config.llm_config.concept_state_names = CONCEPT_STATE_NAMES model_config.llm_config.concept_state_chunk_source_idx = CONCEPT_STATE_CHUNK_SOURCE - model_config.llm_config.concept_state_last_raw_idx = CONCEPT_STATE_LAST_RAW - model_config.llm_config.concept_state_last_final_idx = CONCEPT_STATE_LAST_FINAL + model_config.llm_config.concept_state_last_idx = CONCEPT_STATE_LAST return model_config @staticmethod diff --git a/lmdeploy/pytorch/kernels/cuda/conceptlm.py b/lmdeploy/pytorch/kernels/cuda/conceptlm.py index fb53178dbb..b329b26a1d 100644 --- a/lmdeploy/pytorch/kernels/cuda/conceptlm.py +++ b/lmdeploy/pytorch/kernels/cuda/conceptlm.py @@ -13,7 +13,6 @@ def _decode_chunk_state_update_kernel( state_ids, position_ids, concept_inputs, - next_rows, update_mask, state_stride_n, state_stride_s, @@ -24,9 +23,6 @@ def _decode_chunk_state_update_kernel( out_stride_b, out_stride_s, out_stride_h, - next_stride_b, - next_stride_s, - next_stride_h, HIDDEN: tl.constexpr, TOTAL_ELEMS: tl.constexpr, CHUNK_SIZE: tl.constexpr, @@ -69,12 +65,9 @@ def _decode_chunk_state_update_kernel( zero = tl.zeros((BLOCK, ), dtype=tl.float32) next_value = tl.where(is_boundary, zero, merged) concept_value = tl.where(valid_state & is_boundary, concept, zero) - next_debug_value = tl.where(valid_state, next_value, previous) concept_ptrs = concept_inputs + batch_id * out_stride_b + source_id * out_stride_s + hidden_id * out_stride_h - next_ptrs = next_rows + batch_id * next_stride_b + source_id * next_stride_s + hidden_id * next_stride_h tl.store(concept_ptrs, concept_value, mask=valid_elem) - tl.store(next_ptrs, next_debug_value, mask=valid_elem) tl.store(state_ptrs, next_value, mask=valid_elem & valid_state) if tile_id == 0: @@ -287,10 +280,9 @@ def decode_chunk_state_update( block: Triton vector width. Returns: - Tuple ``(concept_inputs, next_rows, update_mask)``. ``concept_inputs`` - is zero for non-boundary rows. ``next_rows`` is a debug/reference copy - of the per-batch rows written to state cache. ``update_mask`` is - ``True`` only for valid boundary rows. + Tuple ``(concept_inputs, update_mask)``. ``concept_inputs`` is zero for + non-boundary rows. ``update_mask`` is ``True`` only for valid boundary + rows. """ assert chunk_source_state_cache.is_cuda, 'ConceptLM chunk-state kernel requires CUDA state cache.' assert current_source_states.is_cuda, 'ConceptLM chunk-state kernel requires CUDA current states.' @@ -303,7 +295,6 @@ def decode_chunk_state_update( state_ids = state_ids.to(device=current_source_states.device, dtype=torch.long) position_ids = _flatten_decode_position_ids(position_ids, batch_size, current_source_states.device) concept_inputs = torch.empty_like(current_source_states) - next_rows = torch.empty_like(current_source_states) update_mask = torch.empty((batch_size, ), dtype=torch.bool, device=current_source_states.device) grid = (triton.cdiv(total_elems, block), batch_size) _decode_chunk_state_update_kernel[grid]( @@ -312,12 +303,10 @@ def decode_chunk_state_update( state_ids, position_ids, concept_inputs, - next_rows, update_mask, *chunk_source_state_cache.stride(), *current_source_states.stride(), *concept_inputs.stride(), - *next_rows.stride(), HIDDEN=hidden, TOTAL_ELEMS=total_elems, CHUNK_SIZE=int(chunk_size), @@ -325,7 +314,7 @@ def decode_chunk_state_update( BLOCK=block, num_warps=8, ) - return concept_inputs, next_rows, update_mask + return concept_inputs, update_mask def decode_kv_cache_snapshot( diff --git a/lmdeploy/pytorch/models/intern_ncp.py b/lmdeploy/pytorch/models/intern_ncp.py deleted file mode 100644 index 6a4f30a3be..0000000000 --- a/lmdeploy/pytorch/models/intern_ncp.py +++ /dev/null @@ -1,2697 +0,0 @@ -# Copyright (c) OpenMMLab. All rights reserved. -from collections.abc import Iterable -from dataclasses import dataclass, replace -from typing import Any - -import torch -from torch import nn -from torch.nn import functional as F -from transformers.configuration_utils import PretrainedConfig - -from lmdeploy.pytorch.model_inputs import StepContext, StepContextManager -from lmdeploy.pytorch.nn import ( - ApplyRotaryEmb, - Attention, - ConceptLMRuntimeOps, - RMSNorm, - SiluAndMul, - build_rotary_embedding, -) -from lmdeploy.pytorch.nn.linear import ( - build_down_linear, - build_gateup_linear, - build_merged_colwise_linear, - build_o_proj, - build_qkv_proj, -) -from lmdeploy.pytorch.weight_loader.model_weight_loader import load_weight - -from .patch import add_prefix -from .utils.cudagraph import CudaGraphMixin -from .utils.model import DeployModelMixinV1, build_embedding - -_CONFIG_VALUE = object() -_HistoryStates = list[torch.Tensor] | tuple[torch.Tensor, ...] | torch.Tensor -_SourceStates = list[torch.Tensor] | tuple[torch.Tensor, ...] | torch.Tensor | None - - -@dataclass -class ConceptMetadata: - """Layer-invariant ConceptLM runtime metadata. - - This mirrors the DSV4 pattern: ``StepContext`` is read at the top-level - model boundary, then submodules receive explicit metadata instead of - reaching back into the engine context. The dense/reference helpers below do - not consume all fields yet; they are part of the serving decode contract. - """ - - chunk_size: int - merge_method: str - shift_feature: bool - is_decoding: bool | None = None - state_ids: torch.Tensor | None = None - position_ids: torch.Tensor | None = None - block_offsets: torch.Tensor | None = None - q_seqlens: torch.Tensor | None = None - kv_seqlens: torch.Tensor | None = None - q_start_loc: torch.Tensor | None = None - attn_metadata: Any = None - - @classmethod - def build(cls, - config: PretrainedConfig, - position_ids: torch.Tensor | None = None, - attn_metadata: Any = None, - state_ids: torch.Tensor | None = None): - """Build ConceptLM metadata from explicit forward inputs.""" - return cls( - chunk_size=int(config.concept_chunk_size), - merge_method=getattr(config, 'concept_chunk_merge_method', 'meanpooling'), - shift_feature=bool(getattr(config, 'concept_shift_feature', True)), - is_decoding=getattr(attn_metadata, 'is_decoding', None), - state_ids=state_ids, - position_ids=position_ids, - block_offsets=getattr(attn_metadata, 'block_offsets', None), - q_seqlens=getattr(attn_metadata, 'q_seqlens', None), - kv_seqlens=getattr(attn_metadata, 'kv_seqlens', None), - q_start_loc=getattr(attn_metadata, 'q_start_loc', None), - attn_metadata=attn_metadata, - ) - - -@dataclass -class ConceptCaches: - """ConceptLM cache views resolved once at the top-level model boundary.""" - - encoder_past_key_values: list[list[torch.Tensor]] | None = None - concept_past_key_values: list[list[torch.Tensor]] | None = None - decoder_past_key_values: list[list[torch.Tensor]] | None = None - state_caches: list[torch.Tensor] | None = None - chunk_source_idx: int = 0 - last_raw_idx: int = 1 - last_final_idx: int = 2 - - @classmethod - def build(cls, - config: PretrainedConfig, - past_key_values: list[list[torch.Tensor]] | None = None, - state_caches: list[torch.Tensor] | None = None): - """Build ConceptLM cache views from engine-provided caches.""" - encoder_past_key_values, concept_past_key_values, decoder_past_key_values = ( - _split_concept_past_key_values(config, past_key_values)) - return cls( - encoder_past_key_values=encoder_past_key_values, - concept_past_key_values=concept_past_key_values, - decoder_past_key_values=decoder_past_key_values, - state_caches=state_caches, - chunk_source_idx=int(getattr(config, 'concept_state_chunk_source_idx', 0)), - last_raw_idx=int(getattr(config, 'concept_state_last_raw_idx', 1)), - last_final_idx=int(getattr(config, 'concept_state_last_final_idx', 2)), - ) - - def state_cache(self, state_idx: int) -> torch.Tensor | None: - """Return one anonymous state-cache tensor by semantic index.""" - if self.state_caches is None: - return None - if state_idx < 0 or state_idx >= len(self.state_caches): - return None - return self.state_caches[state_idx] - - @property - def chunk_source_state(self) -> torch.Tensor | None: - """Current chunk source accumulator state cache.""" - return self.state_cache(self.chunk_source_idx) - - @property - def last_raw_states(self) -> torch.Tensor | None: - """Latest raw concept-layer state cache.""" - return self.state_cache(self.last_raw_idx) - - @property - def last_final_state(self) -> torch.Tensor | None: - """Latest final concept vector state cache.""" - return self.state_cache(self.last_final_idx) - - -@dataclass -class ConceptChunkStateUpdateResult: - """Fixed-shape result of one decode chunk-source state update.""" - - concept_input_states: torch.Tensor - next_chunk_source_states: torch.Tensor - concept_update_mask: torch.Tensor - valid_state_mask: torch.Tensor - state_ids: torch.Tensor - safe_state_ids: torch.Tensor - - -@dataclass -class ConceptDecodeMetadata: - """Fixed-layout decode metadata derived once from engine inputs. - - Decode is always represented as the engine's fixed ``[1, batch]`` token - layout at the model boundary and flattened to ``[batch]`` / ``[batch, H]`` - only inside ConceptLM helpers. - """ - - position_ids: torch.Tensor - state_ids: torch.Tensor - safe_state_ids: torch.Tensor - valid_state_mask: torch.Tensor - - -@dataclass -class ConceptPrefillMetadata: - """Packed prefill metadata derived once from token attention metadata. - - Field groups: - - token stream: original engine-provided token request boundaries. - - concept stream: compact chunk-token request boundaries and positions - used by the concept predictor attention. - - chunk merge: token -> concept ids and helper ids used to reduce encoder - token states into concept states without a Python batch loop. - - repeat/gather: concept -> token ids used to project compact concept - states back to the packed token stream. - - scalar bounds: eager compact sizes / upper bounds needed by metadata and - attention launch parameters. - """ - - # Token stream metadata, shape [batch]. This is the original packed prefill - # layout consumed by normal token attention. - token_q_seqlens: torch.Tensor - token_q_start_loc: torch.Tensor - - # Concept stream metadata, shape [batch] plus compact concept positions. - # These describe the shorter chunk-token stream consumed by concept - # predictor attention. - concept_q_seqlens: torch.Tensor - concept_q_start_loc: torch.Tensor - concept_position_ids: torch.Tensor - - # Chunk merge metadata. ``merge_token_to_concept`` maps each packed token to - # the compact concept row that owns it, or -1 when the token is dropped from - # concept production. Counts/first/last ids implement mean/first/last merge. - merge_token_to_concept: torch.Tensor - merge_token_counts: torch.Tensor - merge_first_token_ids: torch.Tensor - merge_last_token_ids: torch.Tensor - merge_short_concept_mask: torch.Tensor - - # Repeat/gather metadata. Maps each packed token row to the compact concept - # row it should read after shift semantics are applied, or -1 for the - # zero-concept row. - token_to_concept: torch.Tensor - - # Scalar sizes/bounds. ``num_concepts_total`` is the exact compact size in - # eager prefill; ``max_concepts_per_request`` is the per-request attention - # launch bound. - num_tokens_total: int - num_concepts_total: int - max_concepts_per_request: int - - -def _get_configured_window(config: PretrainedConfig): - """Return the reference OLMo window setting as an int or None.""" - window_size = getattr(config, 'window_size', None) - if window_size is None: - return None - if isinstance(window_size, (list, tuple)): - window_size = window_size[0] - if window_size is None: - return None - window_size = int(window_size) - return window_size if window_size > 0 else None - - -def _make_olmo_rotary_embedding(config: PretrainedConfig, - device: torch.device = None) -> nn.Module: - """Build ConceptLM/OLMo rotary embedding from its native config fields.""" - rotary_interleaved = bool(getattr(config, 'rotary_interleaved', False)) - if rotary_interleaved: - raise NotImplementedError('ConceptLM rotary_interleaved=True is not supported by the LMDeploy block yet.') - - head_dim = int(config.kv_channels) - rotary_percent = float(getattr(config, 'rotary_percent', 1.0)) - rotary_dim = int(head_dim * rotary_percent) - rotary_dim -= rotary_dim % 2 - if rotary_dim <= 0: - raise ValueError(f'Invalid ConceptLM rotary dimension: head_dim={head_dim}, rotary_percent={rotary_percent}') - - partial_rotary_factor = rotary_dim / head_dim - return build_rotary_embedding( - dim=head_dim, - max_position_embeddings=getattr(config, 'max_position_embeddings', getattr(config, 'max_sequence_length', - 2048)), - base=getattr(config, 'rotary_base', 10000), - partial_rotary_factor=partial_rotary_factor, - device=device, - ) - - -def _repack_olmo_qkv_weight(loaded_weight: torch.Tensor, num_heads: int, head_dim: int): - """Convert native OLMo per-head [Q,K,V] QKV packing to LMDeploy - [Q][K][V].""" - leading_shape = loaded_weight.shape[1:] - loaded_weight = loaded_weight.reshape(num_heads, 3, head_dim, *leading_shape) - query = loaded_weight[:, 0].flatten(0, 1) - key = loaded_weight[:, 1].flatten(0, 1) - value = loaded_weight[:, 2].flatten(0, 1) - return torch.cat([query, key, value], dim=0) - - -def _load_stacked_codebook_weight(param: torch.nn.Parameter, loaded_weight: torch.Tensor, codebook_idx: int): - """Load one native ``codebook.N`` checkpoint tensor into a stacked - codebook.""" - assert 0 <= codebook_idx < param.size(0), f'Invalid codebook index: {codebook_idx}' - target = param.data[codebook_idx] - assert target.size() == loaded_weight.size(), ( - f'Attempted to load codebook weight ({loaded_weight.size()}) into parameter slice ({target.size()})') - target.copy_(loaded_weight) - - -def _split_concept_past_key_values(config: PretrainedConfig, past_key_values: list[list[torch.Tensor]] | None): - """Split the flat LMDeploy KV-cache list into ConceptLM streams.""" - if past_key_values is None or len(past_key_values) == 0: - return None, None, None - enc_layers = int(config.concept_encoder_layers) - concept_layers = int(config.concept_special_layers) - dec_layers = int(config.concept_decoder_layers) - total_layers = enc_layers + concept_layers + dec_layers - assert len(past_key_values) >= total_layers, ( - f'ConceptLM requires {total_layers} KV-cache layers ' - f'({enc_layers} encoder + {concept_layers} concept + {dec_layers} decoder), ' - f'got {len(past_key_values)}.') - enc_end = enc_layers - concept_end = enc_end + concept_layers - dec_end = concept_end + dec_layers - return past_key_values[:enc_end], past_key_values[enc_end:concept_end], past_key_values[concept_end:dec_end] - - -def _flatten_decode_position_ids(position_ids: torch.Tensor, batch_size: int) -> torch.Tensor: - """Normalize decode position ids to one absolute position per batch row.""" - if position_ids.dim() == 0: - position_ids = position_ids.view(1) - if position_ids.dim() == 1: - return position_ids.to(torch.long) - position_ids = position_ids.reshape(-1) - if position_ids.numel() == batch_size: - return position_ids.to(torch.long) - assert position_ids.numel() % batch_size == 0, ( - f'Cannot map position_ids with {position_ids.numel()} elements to batch size {batch_size}.') - return position_ids.reshape(-1, batch_size)[-1].to(torch.long) - - -def _concept_decode_chunk_state_update( - chunk_source_state_cache: torch.Tensor, - current_source_states: torch.Tensor, - state_ids: torch.Tensor, - position_ids: torch.Tensor, - chunk_size: int, - merge_method: str, -) -> ConceptChunkStateUpdateResult: - """Compute one decode step's chunk-source state update. - - This helper is intentionally fixed-shape over the decode batch. It returns - per-row next states plus a device-side boundary mask; it does not compact - concept rows by ``num_concepts_total``. The future Triton/CUDA op should - fuse this compute with the state write and skip ``state_id < 0`` rows. - - Args: - chunk_source_state_cache: ``[num_state_slots, num_sources, hidden]``. - current_source_states: ``[batch, num_sources, hidden]`` for the current - decode token after encoder source selection. These should be the - unnormalized states that are merged over the current concept chunk. - state_ids: ``[batch]`` state-cache slot per row, with ``-1`` for - padded CUDA-graph rows. - position_ids: absolute token positions for the decode rows. - chunk_size: ConceptLM chunk size. - merge_method: ``meanpooling``, ``first``, or ``last``. - """ - assert current_source_states.dim() == 3, ( - f'current_source_states must be [batch, num_sources, hidden], got {tuple(current_source_states.shape)}.') - assert chunk_source_state_cache.dim() == 3, ( - f'chunk_source_state_cache must be [num_state_slots, num_sources, hidden], ' - f'got {tuple(chunk_source_state_cache.shape)}.') - batch_size = current_source_states.size(0) - assert current_source_states.shape[1:] == chunk_source_state_cache.shape[1:], ( - f'Current source state shape {tuple(current_source_states.shape[1:])} does not match state-cache shape ' - f'{tuple(chunk_source_state_cache.shape[1:])}.') - - state_ids = state_ids.to(device=current_source_states.device, dtype=torch.long) - position_ids = _flatten_decode_position_ids(position_ids, batch_size).to(device=current_source_states.device) - assert position_ids.numel() == batch_size, ( - f'Expected {batch_size} decode position ids, got {position_ids.numel()}.') - - valid_state_mask = state_ids >= 0 - safe_state_ids = state_ids.clamp(min=0) - previous_rows = chunk_source_state_cache.index_select(0, safe_state_ids) - - chunk_size = int(chunk_size) - chunk_pos = torch.remainder(position_ids, chunk_size) - concept_update_mask = valid_state_mask & (torch.remainder(position_ids + 1, chunk_size) == 0) - first_token_mask = valid_state_mask & (chunk_pos == 0) - merge_method = str(merge_method) - - if merge_method == 'first': - update_rows = torch.where(first_token_mask.view(batch_size, 1, 1), current_source_states, previous_rows) - concept_input_states = update_rows - elif merge_method == 'last': - update_rows = current_source_states - concept_input_states = current_source_states - else: - update_rows = previous_rows + current_source_states - concept_input_states = update_rows / chunk_size - - zero_rows = torch.zeros_like(update_rows) - next_rows = torch.where(concept_update_mask.view(batch_size, 1, 1), zero_rows, update_rows) - next_rows = torch.where(valid_state_mask.view(batch_size, 1, 1), next_rows, previous_rows) - concept_input_states = torch.where(concept_update_mask.view(batch_size, 1, 1), concept_input_states, zero_rows) - return ConceptChunkStateUpdateResult( - concept_input_states=concept_input_states, - next_chunk_source_states=next_rows, - concept_update_mask=concept_update_mask, - valid_state_mask=valid_state_mask, - state_ids=state_ids, - safe_state_ids=safe_state_ids, - ) - - -def _apply_concept_chunk_state_update_reference_(chunk_source_state_cache: torch.Tensor, - update: ConceptChunkStateUpdateResult): - """Reference-only state write for tests. - - This uses a Python loop and may read scalar state ids on host. Do not call - it from the serving hot path. The graph-safe implementation should write - ``update.next_chunk_source_states`` inside a backend op that skips - ``state_id < 0`` rows. - """ - for batch_idx in range(update.state_ids.numel()): - state_id = int(update.state_ids[batch_idx]) - if state_id < 0: - continue - chunk_source_state_cache[state_id].copy_(update.next_chunk_source_states[batch_idx]) - - -def _qk_rmsnorm_variance(query: torch.Tensor, key: torch.Tensor) -> torch.Tensor: - """Local Q/K squared sums before TP all-reduce.""" - query = query.float() - key = key.float() - query_var = (query * query).sum(-1, keepdim=True) - key_var = (key * key).sum(-1, keepdim=True) - return torch.stack([query_var, key_var], dim=0) - - -def _qk_rmsnorm_apply(query: torch.Tensor, - key: torch.Tensor, - variance: torch.Tensor, - query_weight: torch.Tensor, - key_weight: torch.Tensor, - hidden_size: int, - eps: float): - """Apply whole-hidden Q/K RMSNorm from an already all-reduced variance.""" - dtype = query.dtype - query_var, key_var = variance / hidden_size + eps - query = (query.float() * torch.rsqrt(query_var)).to(dtype) * query_weight - key = (key.float() * torch.rsqrt(key_var)).to(dtype) * key_weight - return query, key - - -class ConceptLMV22VQEmbedding(nn.Module): - """Token embedding container. - - Mirrors ``self.embedding`` in the reference: a plain ``nn.Module`` holding - a ``word_embeddings`` table. Keeping the attribute name ``word_embeddings`` - makes the checkpoint key ``embedding.word_embeddings.weight`` line up - directly with ``self.embedding.word_embeddings``. - """ - - def __init__(self, - config: PretrainedConfig, - dtype: torch.dtype = None, - device: torch.device = None): - super().__init__() - self.word_embeddings = build_embedding( - config.vocab_size, - config.hidden_size, - getattr(config, 'pad_token_id', None), - dtype=dtype, - device=device, - ) - - def forward(self, input_ids: torch.Tensor) -> torch.Tensor: - """forward.""" - return self.word_embeddings(input_ids) - - -class ConceptLMV22VQQuantizer(nn.Module): - """Rewrite of ``_Quantizer``. - - The reference keeps codebooks as a ``ParameterList`` only to produce - checkpoint keys ``codebook.0`` ... ``codebook.N``. Runtime computation only - needs the stacked tensor returned by ``transformed_codebook()``. - """ - - def __init__(self, - config: PretrainedConfig, - dtype: torch.dtype = None, - device: torch.device = None): - super().__init__() - self.num_codebooks = int(config.concept_v22_vq_num_codebooks) - self.codebook_size = int(config.concept_v22_vq_codebook_size) - hidden_size = int(config.hidden_size) - assert hidden_size % self.num_codebooks == 0, ( - f'hidden_size={hidden_size} must be divisible by num_codebooks={self.num_codebooks}') - self.codebook_dim = hidden_size // self.num_codebooks - self.hidden_size = hidden_size - self.codebook = nn.Parameter( - torch.empty( - self.num_codebooks, - self.codebook_size, - self.codebook_dim, - dtype=dtype, - device=device, - ), - requires_grad=False, - ) - - def transformed_codebook(self): - """Return codebook as ``[num_codebooks, codebook_size, - codebook_dim]``.""" - return self.codebook - - def forward(self, concept_logits: torch.Tensor) -> torch.Tensor: - """Convert per-codebook logits to hidden vectors. - - Args: - concept_logits: ``[..., num_codebooks, codebook_size]``. In the - LMDeploy engine the leading dims may be a packed continuous - batching dimension, e.g. ``[num_concepts_total]``. - - Returns: - Quantized/predicted vectors with shape ``[..., hidden_size]``. - """ - assert concept_logits.shape[-2:] == (self.num_codebooks, self.codebook_size), ( - f'Expected concept logits trailing shape {(self.num_codebooks, self.codebook_size)}, ' - f'got {tuple(concept_logits.shape[-2:])}.') - codebook = self.transformed_codebook().to(concept_logits.dtype) - vectors = torch.einsum('...hk,hkd->...hd', concept_logits, codebook) - return vectors.flatten(-2, -1) - - -class ConceptLMV22VQDepthDD(nn.Module): - """Rewrite of ``_DepthDD``. - - This is a small replicated per-token depth mixer. It does not mix tokens; - it computes ``num_prev`` route weights from the current hidden state and - combines the matching per-layer history states. The reference only handles - dense ``[seq, batch, hidden]`` tensors because it stacks history at - ``dim=2``. LMDeploy continuous batching commonly uses packed - ``[num_tokens, hidden]`` tensors. Runtime should pass a preallocated tensor - history to avoid repeated ``torch.stack`` copies; list/tuple input remains - only as a debug/parity convenience. - """ - - def __init__(self, - config: PretrainedConfig, - layer_idx: int, - use_softmax: bool, - dtype: torch.dtype = None, - device: torch.device = None): - super().__init__() - self.num_prev = int(layer_idx) + 2 - route_hidden_size = self.num_prev - self.w1 = nn.Linear(config.hidden_size, route_hidden_size, bias=False, dtype=dtype, device=device) - self.w2 = nn.Linear(route_hidden_size, self.num_prev, bias=False, dtype=dtype, device=device) - self.static_a = nn.Parameter(torch.zeros(self.num_prev, dtype=dtype, device=device), requires_grad=False) - self.use_softmax = bool(use_softmax) - for param in self.parameters(): - param.requires_grad_(False) - - def _history_tensor(self, hidden_states: torch.Tensor, history_states: _HistoryStates, history_dim: int): - """Return history in shape ``[..., num_prev, hidden_size]``.""" - if not isinstance(history_states, torch.Tensor): - assert len(history_states) == self.num_prev, ( - f'Expected {self.num_prev} history states, got {len(history_states)}.') - history_states = torch.stack(tuple(history_states), dim=-2) - history_dim = -2 - - history_dim = history_dim if history_dim >= 0 else history_dim + history_states.dim() - assert 0 <= history_dim < history_states.dim(), f'Invalid history_dim={history_dim}.' - assert history_states.shape[history_dim] == self.num_prev, ( - f'Expected history dimension {history_dim} to be {self.num_prev}, ' - f'got {history_states.shape[history_dim]}.') - if history_dim != history_states.dim() - 2: - history_states = history_states.movedim(history_dim, -2) - assert history_states.shape[:-2] == hidden_states.shape[:-1], ( - f'Expected history leading shape {tuple(hidden_states.shape[:-1])}, ' - f'got {tuple(history_states.shape[:-2])}.') - assert history_states.shape[-1] == hidden_states.shape[-1], ( - f'Expected history hidden size {hidden_states.shape[-1]}, got {history_states.shape[-1]}.') - return history_states - - def forward(self, - hidden_states: torch.Tensor, - history_states: _HistoryStates, - history_dim: int = -2) -> torch.Tensor: - """forward.""" - history = self._history_tensor(hidden_states, history_states, history_dim) - weights = self.w2(F.gelu(self.w1(hidden_states))) - weights = weights + self.static_a.to(dtype=weights.dtype) - if self.use_softmax: - weights = weights.softmax(dim=-1) - return torch.einsum('...l,...lh->...h', weights, history) - - -class ConceptLMV22VQSelfDD(nn.Module): - """Rewrite of ``_SelfDD``.""" - - def __init__(self, - config: PretrainedConfig, - num_layers: int, - use_softmax: bool = False, - dtype: torch.dtype = None, - device: torch.device = None): - super().__init__() - self.num_layers = int(num_layers) - self.depth_dds = nn.ModuleList([ - ConceptLMV22VQDepthDD(config, layer_idx, use_softmax, dtype=dtype, device=device) - for layer_idx in range(self.num_layers) - ]) - - def make_history_buffer(self, hidden_states: torch.Tensor) -> torch.Tensor: - """Allocate layer-major history buffer ``[num_layers + 1, - *hidden_shape]``.""" - return hidden_states.new_empty((self.num_layers + 1, *hidden_states.shape)) - - @staticmethod - def write_history(history_buffer: torch.Tensor, slot_idx: int, hidden_states: torch.Tensor): - """Copy one history block into a layer-major history buffer.""" - assert history_buffer.dim() == hidden_states.dim() + 1, ( - f'Expected history buffer dim {hidden_states.dim() + 1}, got {history_buffer.dim()}.') - assert history_buffer.shape[1:] == hidden_states.shape, ( - f'Expected history buffer trailing shape {tuple(hidden_states.shape)}, ' - f'got {tuple(history_buffer.shape[1:])}.') - history_buffer[int(slot_idx)].copy_(hidden_states) - return history_buffer - - @staticmethod - def history_view(history_buffer: torch.Tensor, layer_idx: int): - """Return layer-major history needed by ``layer_idx`` without copying. - - This is CUDA-graph safe when ``layer_idx`` is a Python constant for the - current layer and ``history_buffer`` has the fixed graph-capture shape. - """ - return history_buffer[:int(layer_idx) + 2] - - def forward_from_buffer(self, - layer_idx: int, - hidden_states: torch.Tensor, - history_buffer: torch.Tensor): - """Runtime path: fixed buffer, no list materialization.""" - layer_idx = int(layer_idx) - return self.depth_dds[layer_idx]( - hidden_states, - self.history_view(history_buffer, layer_idx), - history_dim=0, - ) - - def forward(self, - layer_idx: int, - hidden_states: torch.Tensor, - history_states: _HistoryStates): - """forward.""" - return self.depth_dds[int(layer_idx)](hidden_states, history_states) - - -class ConceptLMV22VQResidualRoute(nn.Module): - """Rewrite of ``_ResidualRoute``. - - This module computes a source-state mixture and adds it as a gated residual - update to the target hidden state. It is small and replicated. - - Runtime/static-graph code should use ``forward_from_buffer`` with a full - fixed-size source buffer. Flexible active-source/list paths are retained for - parity tests and WIP reference wiring, but they should not be used inside a - captured CUDA graph because they can change intermediate shapes. - """ - - def __init__(self, - config: PretrainedConfig, - num_source_states: int, - use_softmax: bool = True, - dtype: torch.dtype = None, - device: torch.device = None): - super().__init__() - self.num_source_states = int(num_source_states) - route_hidden_size = max(1, self.num_source_states) - self.w1 = nn.Linear(config.hidden_size, route_hidden_size, bias=False, dtype=dtype, device=device) - self.w2 = nn.Linear(route_hidden_size, self.num_source_states, bias=False, dtype=dtype, device=device) - self.residual_diag = nn.Parameter(torch.zeros(config.hidden_size, dtype=dtype, device=device), - requires_grad=False) - self.use_softmax = bool(use_softmax) - for param in self.parameters(): - param.requires_grad_(False) - - def _source_tensor(self, - target_hidden: torch.Tensor, - source_states: _SourceStates, - source_dim: int, - expected_leading_shape: tuple[int, ...] | None = None): - """Return source states in shape ``[..., active_sources, - hidden_size]``.""" - if source_states is None: - return None - if not isinstance(source_states, torch.Tensor): - if len(source_states) == 0: - return None - source_states = torch.stack(tuple(source_states), dim=-2) - source_dim = -2 - - source_dim = source_dim if source_dim >= 0 else source_dim + source_states.dim() - assert 0 <= source_dim < source_states.dim(), f'Invalid source_dim={source_dim}.' - if source_states.shape[source_dim] == 0: - return None - if source_dim != source_states.dim() - 2: - source_states = source_states.movedim(source_dim, -2) - if expected_leading_shape is None: - expected_leading_shape = target_hidden.shape[:-1] - assert source_states.shape[:-2] == expected_leading_shape, ( - f'Expected source leading shape {tuple(expected_leading_shape)}, ' - f'got {tuple(source_states.shape[:-2])}.') - assert source_states.shape[-1] == target_hidden.shape[-1], ( - f'Expected source hidden size {target_hidden.shape[-1]}, got {source_states.shape[-1]}.') - return source_states - - def _route_weights(self, target_hidden: torch.Tensor, active_sources: int): - """Compute route weights and keep the last active source logits.""" - weights = self.w2(F.gelu(self.w1(target_hidden))) - weights = weights[..., -active_sources:] - if self.use_softmax: - weights = weights.softmax(dim=-1) - return weights - - def _add_update(self, - target_hidden: torch.Tensor, - source_mix: torch.Tensor, - residual_scale: torch.Tensor | None = None): - """Apply residual diagonal and optional scale, then add to target.""" - update = source_mix * self.residual_diag.to(source_mix.dtype) - if residual_scale is not None: - update = update * residual_scale.to(update.dtype) - return target_hidden + update.to(target_hidden.dtype) - - def make_source_buffer(self, hidden_states: torch.Tensor) -> torch.Tensor: - """Allocate source-major buffer ``[num_source_states, - *hidden_shape]``.""" - return hidden_states.new_empty((self.num_source_states, *hidden_states.shape)) - - @staticmethod - def write_source(source_buffer: torch.Tensor, slot_idx: int, source_state: torch.Tensor): - """Copy one source state into a source-major buffer.""" - assert source_buffer.dim() == source_state.dim() + 1, ( - f'Expected source buffer dim {source_state.dim() + 1}, got {source_buffer.dim()}.') - assert source_buffer.shape[1:] == source_state.shape, ( - f'Expected source buffer trailing shape {tuple(source_state.shape)}, ' - f'got {tuple(source_buffer.shape[1:])}.') - source_buffer[int(slot_idx)].copy_(source_state) - return source_buffer - - @staticmethod - def source_view(source_buffer: torch.Tensor, active_sources: int | None = None): - """Return active source-major view without copying. - - Passing ``active_sources`` is a flexible/debug path. Static graph - runtime should pass full fixed buffers and leave ``active_sources`` as - ``None``. - """ - if active_sources is None: - return source_buffer - return source_buffer[:int(active_sources)] - - def forward(self, - target_hidden: torch.Tensor, - source_states: _SourceStates, - residual_scale: torch.Tensor | None = None, - source_dim: int = -2): - """Flexible/debug forward path. - - This accepts lists, ``None``, and active source tensors for reference - parity. Use ``forward_from_buffer`` for static-shape runtime. - """ - source_states = self._source_tensor(target_hidden, source_states, source_dim) - if source_states is None: - return target_hidden - active_sources = source_states.shape[-2] - weights = self._route_weights(target_hidden, active_sources) - source_mix = torch.einsum('...m,...mh->...h', weights, source_states) - return self._add_update(target_hidden, source_mix, residual_scale) - - def forward_from_buffer(self, - target_hidden: torch.Tensor, - source_buffer: torch.Tensor, - residual_scale: torch.Tensor | None = None): - """Runtime path: read a full fixed source-major buffer. - - ``source_buffer`` shape is ``[num_source_states, *target_hidden.shape]``. - This keeps source count fixed across CUDA graph capture/replay. - """ - assert source_buffer.shape[0] == self.num_source_states, ( - f'Expected full source buffer with {self.num_source_states} states, got {source_buffer.shape[0]}.') - return self.forward( - target_hidden, - source_buffer, - residual_scale=residual_scale, - source_dim=0, - ) - - def forward_active_from_buffer(self, - target_hidden: torch.Tensor, - source_buffer: torch.Tensor, - active_sources: int, - residual_scale: torch.Tensor | None = None): - """Flexible/debug buffer path with reduced active source count.""" - return self.forward( - target_hidden, - self.source_view(source_buffer, active_sources), - residual_scale=residual_scale, - source_dim=0, - ) - - def forward_repeated_chunks(self, - target_hidden: torch.Tensor, - source_states: _SourceStates, - chunk_size: int, - shift_feature: bool, - residual_scale: torch.Tensor | None = None, - source_dim: int = 2): - """Reference repeated-chunk route used by decoder-read-concept. - - ``target_hidden`` is dense ``[seq, batch, hidden]`` and source states are - chunk-level ``[chunks, batch, sources, hidden]`` after ``source_dim`` is - normalized to 2. A packed continuous-batching runtime will need token to - chunk metadata before using this path end-to-end. This method is not a - complete CUDA-graph runtime path yet because chunk lengths still need a - fixed-buffer/mask contract at the caller level. - """ - if source_states is None: - return target_hidden - assert target_hidden.dim() == 3, ( - f'forward_repeated_chunks expects [seq, batch, hidden], got {tuple(target_hidden.shape)}.') - - if not isinstance(source_states, torch.Tensor): - if len(source_states) == 0: - return target_hidden - source_states = torch.stack(tuple(source_states), dim=2) - source_dim = 2 - source_dim = source_dim if source_dim >= 0 else source_dim + source_states.dim() - assert 0 <= source_dim < source_states.dim(), f'Invalid source_dim={source_dim}.' - if source_states.shape[source_dim] == 0: - return target_hidden - if source_dim != 2: - source_states = source_states.movedim(source_dim, 2) - - seq_len, batch_size, hidden_size = target_hidden.shape - assert source_states.dim() == 4, ( - f'Expected source states [chunks, batch, sources, hidden], got {tuple(source_states.shape)}.') - assert source_states.shape[1] == batch_size, ( - f'Expected source batch size {batch_size}, got {source_states.shape[1]}.') - assert source_states.shape[3] == hidden_size, ( - f'Expected source hidden size {hidden_size}, got {source_states.shape[3]}.') - num_chunks, _, active_sources, _ = source_states.shape - chunk_size = int(chunk_size) - weights = self._route_weights(target_hidden, active_sources) - if shift_feature: - weights = torch.cat((weights.new_zeros(1, batch_size, active_sources), weights), dim=0) - repeated_len = num_chunks * chunk_size - if weights.shape[0] < repeated_len: - pad_len = repeated_len - weights.shape[0] - weights = torch.cat((weights, weights.new_zeros(pad_len, batch_size, active_sources)), dim=0) - weights = weights[:repeated_len] - - source_mix = torch.einsum( - 'ckbm,cbmh->ckbh', - weights.reshape(num_chunks, chunk_size, batch_size, active_sources), - source_states, - ).reshape(repeated_len, batch_size, hidden_size) - if shift_feature: - source_mix = source_mix[1:1 + seq_len] - else: - source_mix = source_mix[:seq_len] - if source_mix.shape[0] < seq_len: - pad_len = seq_len - source_mix.shape[0] - source_mix = torch.cat((source_mix, source_mix.new_zeros(pad_len, batch_size, hidden_size)), dim=0) - return self._add_update(target_hidden, source_mix, residual_scale) - - def forward_repeated_chunks_from_buffer(self, - target_hidden: torch.Tensor, - source_buffer: torch.Tensor, - chunk_size: int, - shift_feature: bool, - residual_scale: torch.Tensor | None = None): - """Full source-major buffer variant of ``forward_repeated_chunks``.""" - assert source_buffer.shape[0] == self.num_source_states, ( - f'Expected full source buffer with {self.num_source_states} states, got {source_buffer.shape[0]}.') - return self.forward_repeated_chunks( - target_hidden, - source_buffer, - chunk_size, - shift_feature, - residual_scale=residual_scale, - source_dim=0, - ) - - def forward_repeated_chunks_active_from_buffer(self, - target_hidden: torch.Tensor, - source_buffer: torch.Tensor, - chunk_size: int, - shift_feature: bool, - active_sources: int, - residual_scale: torch.Tensor | None = None): - """Flexible/debug repeated-chunk buffer path with reduced active source - count.""" - return self.forward_repeated_chunks( - target_hidden, - self.source_view(source_buffer, active_sources), - chunk_size, - shift_feature, - residual_scale=residual_scale, - source_dim=0, - ) - - -class ConceptLMV22VQConceptRoute(nn.Module): - """Rewrite of ``_ConceptRoute``. - - Applies LayerNorm to the final concept state, scales it elementwise with a learned diagonal, optionally applies a - route scale, then adds the update to decoder hidden states. This is replicated and token-local. - """ - - def __init__(self, - config: PretrainedConfig, - dtype: torch.dtype = None, - device: torch.device = None): - super().__init__() - self.concept_norm = nn.LayerNorm(config.hidden_size, - eps=getattr(config, 'layernorm_epsilon', 1e-6), - dtype=dtype, - device=device) - self.final_diag = nn.Parameter(torch.zeros(config.hidden_size, dtype=dtype, device=device), - requires_grad=False) - for param in self.parameters(): - param.requires_grad_(False) - - def forward(self, - hidden_states: torch.Tensor, - final_concept_state: torch.Tensor, - final_scale: torch.Tensor | None = None): - """forward.""" - concept = self.concept_norm(final_concept_state) - update = concept * self.final_diag.to(concept.dtype) - if final_scale is not None: - update = update * final_scale.to(update.dtype) - return hidden_states + update.to(hidden_states.dtype) - - -class ConceptLMV22VQTwoRouteAdd(nn.Module): - """Rewrite of ``_TwoRouteAdd``. - - It first applies decoder-side ``_DepthDD`` to the decoder history, then - injects the final concept state with ``_ConceptRoute``. - """ - - def __init__(self, - config: PretrainedConfig, - dtype: torch.dtype = None, - device: torch.device = None): - super().__init__() - self.num_layers = int(config.concept_decoder_layers) - use_softmax = bool(getattr(config, 'concept_dd_two_route_add_decoder_use_softmax', True)) - self.decoder_dds = nn.ModuleList([ - ConceptLMV22VQDepthDD(config, layer_idx, use_softmax, dtype=dtype, device=device) - for layer_idx in range(self.num_layers) - ]) - self.concept_routes = nn.ModuleList([ - ConceptLMV22VQConceptRoute(config, dtype=dtype, device=device) - for _ in range(self.num_layers) - ]) - - def make_history_buffer(self, hidden_states: torch.Tensor) -> torch.Tensor: - """Allocate layer-major decoder history buffer ``[num_layers + 1, - *hidden_shape]``.""" - return hidden_states.new_empty((self.num_layers + 1, *hidden_states.shape)) - - @staticmethod - def write_history(history_buffer: torch.Tensor, slot_idx: int, hidden_states: torch.Tensor): - """Copy one decoder history block into a layer-major history buffer.""" - return ConceptLMV22VQSelfDD.write_history(history_buffer, slot_idx, hidden_states) - - @staticmethod - def history_view(history_buffer: torch.Tensor, layer_idx: int): - """Return layer-major decoder history needed by ``layer_idx`` without - copying.""" - return ConceptLMV22VQSelfDD.history_view(history_buffer, layer_idx) - - def forward_from_buffer(self, - layer_idx: int, - hidden_states: torch.Tensor, - history_buffer: torch.Tensor, - final_concept_state: torch.Tensor, - final_scale: torch.Tensor | None = None): - """Runtime path: fixed decoder history buffer, no list materialization.""" - layer_idx = int(layer_idx) - hidden_states = self.decoder_dds[layer_idx]( - hidden_states, - self.history_view(history_buffer, layer_idx), - history_dim=0, - ) - return self.concept_routes[layer_idx](hidden_states, final_concept_state, final_scale) - - def forward(self, - layer_idx: int, - hidden_states: torch.Tensor, - history_states: _HistoryStates, - final_concept_state: torch.Tensor, - final_scale: torch.Tensor | None = None): - """Flexible/debug forward path.""" - layer_idx = int(layer_idx) - hidden_states = self.decoder_dds[layer_idx](hidden_states, history_states) - return self.concept_routes[layer_idx](hidden_states, final_concept_state, final_scale) - - -class ConceptLMV22VQPredictionHeads(nn.Module): - """Merged per-codebook prediction heads. - - Native checkpoints store one ``prediction_heads.N`` linear per concept - codebook. Runtime uses one merged projection and loads each native head into - a deterministic output shard through LMDeploy's ``param.weight_loader`` - pattern. - """ - - def __init__(self, - config: PretrainedConfig, - dtype: torch.dtype = None, - device: torch.device = None, - prefix: str = ''): - codebook_size = int(config.concept_v22_vq_codebook_size) - num_codebooks = int(config.concept_v22_vq_num_codebooks) - super().__init__() - self.num_codebooks = num_codebooks - self.codebook_size = codebook_size - quantization_config = getattr(config, 'quantization_config', None) - self.proj = build_merged_colwise_linear( - config.hidden_size, - [codebook_size] * num_codebooks, - bias=True, - dtype=dtype, - device=device, - quant_config=quantization_config, - # Keep concept logits replicated for now. Output TP would make the - # trailing ``[num_codebooks, codebook_size]`` contract sharded and - # needs a defined distributed sampling/gather path first. - is_tp=False, - out_names=list(range(num_codebooks)), - prefix=add_prefix('proj', prefix), - ) - - def forward(self, hidden_states: torch.Tensor): - """Return logits in shape ``[..., num_codebooks, codebook_size]``.""" - logits = self.proj(hidden_states) - return logits.unflatten(-1, (self.num_codebooks, self.codebook_size)) - - -class ConceptLMV22VQConceptPredictor(nn.Module): - """Rewrite of ``_ConceptPredictor``. - - The predictor owns the high-level concept OLMo block, per-codebook prediction heads, concept self-DD, encoder-read - routes, and the shared source LayerNorm used before encoder states are routed into the concept stream. - """ - - def __init__(self, - config: PretrainedConfig, - dtype: torch.dtype = None, - device: torch.device = None, - prefix: str = ''): - super().__init__() - self.num_layers = int(config.concept_special_layers) - self.num_codebooks = int(config.concept_v22_vq_num_codebooks) - self.codebook_size = int(config.concept_v22_vq_codebook_size) - self.hlm_block = ConceptLMV22VQOlmoBlock(config, - self.num_layers, - post_layer_norm=True, - dtype=dtype, - device=device, - prefix=add_prefix('hlm_block', prefix)) - self.prediction_heads = ConceptLMV22VQPredictionHeads(config, - dtype=dtype, - device=device, - prefix=add_prefix('prediction_heads', prefix)) - self.concept_self_dd = ConceptLMV22VQSelfDD(config, - self.num_layers, - use_softmax=False, - dtype=dtype, - device=device) - self.concept_read_encoder_routes = nn.ModuleList([ - ConceptLMV22VQResidualRoute(config, - int(config.concept_encoder_layers) - 1, - use_softmax=True, - dtype=dtype, - device=device) - for _ in range(self.num_layers) - ]) - self.concept_read_encoder_shared_source_norm = nn.LayerNorm( - config.hidden_size, - eps=getattr(config, 'layernorm_epsilon', 1e-6), - dtype=dtype, - device=device) - for param in self.concept_read_encoder_shared_source_norm.parameters(): - param.requires_grad_(False) - - def set_attention_window(self, window_size, skip_frequency): - """Reference-compatible API. - - LMDeploy's rewrite bakes per-layer sliding-window policy into - ``ConceptLMV22VQOlmoBlock`` at construction time, so this is retained as - an explicit no-op for call-site compatibility. - """ - self._window_size = window_size - self._window_skip_frequency = skip_frequency - - def normalize_encoder_concept_states(self, encoder_concept_states: torch.Tensor): - """Apply the shared source norm used before concept-read-encoder - routes.""" - return self.concept_read_encoder_shared_source_norm(encoder_concept_states) - - def predict_logits(self, hidden_states: torch.Tensor): - """Return logits in shape ``[..., num_codebooks, codebook_size]``.""" - return self.prediction_heads(hidden_states) - - def make_history_buffer(self, hidden_states: torch.Tensor) -> torch.Tensor: - """Allocate concept self-DD history buffer.""" - return self.concept_self_dd.make_history_buffer(hidden_states) - - @staticmethod - def write_history(history_buffer: torch.Tensor, slot_idx: int, hidden_states: torch.Tensor): - """Copy one concept history block into a layer-major history buffer.""" - return ConceptLMV22VQSelfDD.write_history(history_buffer, slot_idx, hidden_states) - - def forward(self, - concept_hidden: torch.Tensor, - encoder_concept_states: torch.Tensor, - position_ids: torch.Tensor, - past_key_values: list[list[torch.Tensor]] | None = None, - attn_metadata: Any = None, - encoder_source_dim: int = -2): - """Concept predictor forward. - - This mirrors the reference control flow, but uses LMDeploy's OLMo layer - rewrite and paged attention inputs. Full end-to-end use requires the - caller to provide concept-stream ``past_key_values`` and - ``attn_metadata`` matching the concept token layout. - """ - hidden_states = concept_hidden - history_buffer = self.make_history_buffer(hidden_states) - self.write_history(history_buffer, 0, hidden_states) - raw_states = [] - rotary_pos_emb = self.hlm_block._make_rotary_pos_emb(hidden_states, position_ids) - - for layer_idx, layer in enumerate(self.hlm_block.layers): - pkv = past_key_values[layer_idx] if past_key_values is not None else None - raw = layer( - hidden_states, - rotary_pos_emb=rotary_pos_emb, - past_key_value=pkv, - attn_metadata=attn_metadata, - ) - raw_states.append(raw) - self.write_history(history_buffer, layer_idx + 1, raw) - hidden_states = self.concept_self_dd.forward_from_buffer(layer_idx, raw, history_buffer) - hidden_states = self.concept_read_encoder_routes[layer_idx]( - hidden_states, - encoder_concept_states, - source_dim=encoder_source_dim, - ) - - hidden_states = self.hlm_block.final_layernorm(hidden_states) - logits = self.predict_logits(hidden_states) - return logits, raw_states - - -class ConceptLMV22VQOlmoAttention(nn.Module): - """Rewrite of ``_OlmoSelfAttention``. - - Differences from the reference: - - fused QKV uses lmdeploy's ``build_qkv_proj`` (standard [Q,K,V] packing, - TP-aware). The reference stores QKV per attention head, so the native - checkpoint weight is re-laid-out during ``load_weights``. - - q/k RMSNorm operates on the whole hidden_size (as in the reference), - not per-head. TP>1 needs an all-reduced variance; we follow internvl's - qkv_norm pattern. - - attention runs through lmdeploy's paged ``Attention`` primitive. - """ - - def __init__(self, - config: PretrainedConfig, - dtype: torch.dtype = None, - device: torch.device = None, - sliding_window: int | None = None, - prefix: str = ''): - super().__init__() - quantization_config = getattr(config, 'quantization_config', None) - hidden_size = config.hidden_size - num_heads = config.num_attention_heads - head_dim = config.kv_channels - num_kv_heads = num_heads - assert hidden_size == num_heads * head_dim, ( - f'ConceptLM OLMo attention expects hidden_size == num_heads * head_dim, ' - f'got {hidden_size} != {num_heads} * {head_dim}.') - - self.num_heads = num_heads - self.num_kv_heads = num_kv_heads - self.head_dim = head_dim - self.hidden_size = hidden_size - self.layernorm_epsilon = config.layernorm_epsilon - - # packed qkv - self.qkv_proj = build_qkv_proj( - hidden_size, - num_q_heads=num_heads, - num_kv_heads=num_kv_heads, - head_size=head_dim, - bias=False, - quant_config=quantization_config, - dtype=dtype, - device=device, - prefix=add_prefix('qkv_proj', prefix), - ) - - # q, k norm over the whole hidden_size (NOT per-head). tp=True with - # head_dim alignment so the weight shards correctly under TP. - self.q_layernorm = RMSNorm(hidden_size, - config.layernorm_epsilon, - quant_config=quantization_config, - dtype=dtype, - device=device, - tp=True, - align=head_dim, - prefix=add_prefix('q_layernorm', prefix)) - self.k_layernorm = RMSNorm(hidden_size, - config.layernorm_epsilon, - quant_config=quantization_config, - dtype=dtype, - device=device, - tp=True, - align=head_dim, - prefix=add_prefix('k_layernorm', prefix)) - - # rotary embedding - self.apply_rotary_pos_emb = ApplyRotaryEmb() - - # attention - self.attn_fwd = Attention(num_heads, - head_dim, - num_kv_heads=num_kv_heads, - v_head_size=head_dim, - sliding_window=None if sliding_window is None else int(sliding_window), - device=device) - - # o_proj - self.o_proj = build_o_proj(num_heads * head_dim, - hidden_size, - bias=False, - quant_config=quantization_config, - dtype=dtype, - device=device, - is_tp=True, - prefix=add_prefix('o_proj', prefix)) - - def _qkv_norm(self, q: torch.Tensor, k: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: - """RMSNorm over the whole hidden_size, TP-correct via all-reduce.""" - import lmdeploy.pytorch.distributed as dist - q_shape = q.shape - k_shape = k.shape - q = q.flatten(-2, -1) - k = k.flatten(-2, -1) - - tp, _ = dist.get_tp_world_rank('attn') - if tp == 1: - q = self.q_layernorm(q) - k = self.k_layernorm(k) - return q.view(q_shape), k.view(k_shape) - - # variance is computed over the full hidden_size, so it must be - # all-reduced across TP ranks before normalizing the local shard. - variance = _qk_rmsnorm_variance(q, k) - dist.all_reduce(variance) - q, k = _qk_rmsnorm_apply(q, k, variance, self.q_layernorm.weight, - self.k_layernorm.weight, self.hidden_size, - self.layernorm_epsilon) - return q.view(q_shape), k.view(k_shape) - - def forward(self, - hidden_states: torch.Tensor, - rotary_pos_emb: tuple[torch.Tensor, torch.Tensor], - past_key_value: list[torch.Tensor] | None = None, - attn_metadata: Any = None): - """Rewrite of _OlmoSelfAttention.forward.""" - # qkv proj -> (batch, seq, num_heads, head_dim) each - qkv_states = self.qkv_proj(hidden_states) - qkv_states = qkv_states.flatten(0, -2) # (-1, heads_total, head_dim) - query_states, key_states, value_states = self.qkv_proj.split_qkv(qkv_states) - - # q, k norm (whole hidden_size) - query_states, key_states = self._qkv_norm(query_states, key_states) - - # rotary embedding - cos, sin = rotary_pos_emb - query_states, key_states = self.apply_rotary_pos_emb( - query_states, key_states, cos, sin, inplace=True) - - # attention (paged) - attn_output = self.attn_fwd( - query_states, - key_states, - value_states, - past_key_value[0], - past_key_value[1], - attn_metadata, - k_scales_zeros=None if len(past_key_value) == 2 else past_key_value[2], - v_scales_zeros=None if len(past_key_value) == 2 else past_key_value[3], - inplace=True, - ) - attn_output = attn_output.reshape(*hidden_states.shape[:-1], -1) - - # o proj - attn_output = self.o_proj(attn_output) - return attn_output - - -class ConceptLMV22VQOlmoMLP(nn.Module): - """Rewrite of ``_OlmoMLP`` (SwiGLU).""" - - def __init__(self, - config: PretrainedConfig, - dtype: torch.dtype = None, - device: torch.device = None, - prefix: str = ''): - super().__init__() - quantization_config = getattr(config, 'quantization_config', None) - self.gate_up_proj = build_gateup_linear( - config.hidden_size, - [config.ffn_hidden_size, config.ffn_hidden_size], - bias=False, - dtype=dtype, - device=device, - quant_config=quantization_config, - is_tp=True, - prefix=add_prefix('gate_up_proj', prefix), - ) - self.act_fn = SiluAndMul(inplace=True) - self.down_proj = build_down_linear( - config.ffn_hidden_size, - config.hidden_size, - bias=False, - quant_config=quantization_config, - dtype=dtype, - device=device, - is_tp=True, - prefix=add_prefix('down_proj', prefix), - ) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - """forward.""" - gate_up = self.gate_up_proj(hidden_states) - act = self.act_fn(gate_up) - return self.down_proj(act) - - -class ConceptLMV22VQOlmoLayer(nn.Module): - """Rewrite of ``_OlmoLayer``. - - The reference uses post-norm residuals:: - h = h + post_attention_layernorm(attn(h)) - h = h + post_feedforward_layernorm(mlp(h)) - This block therefore calls RMSNorm without the residual argument and adds - the residual explicitly. - """ - - def __init__(self, - config: PretrainedConfig, - layer_idx: int, - dtype: torch.dtype = None, - device: torch.device = None, - sliding_window: int | None = None, - prefix: str = ''): - super().__init__() - quantization_config = getattr(config, 'quantization_config', None) - self.layer_number = int(layer_idx) + 1 - self.self_attention = ConceptLMV22VQOlmoAttention(config, - dtype=dtype, - device=device, - sliding_window=sliding_window, - prefix=add_prefix('self_attention', prefix)) - self.post_attention_layernorm = RMSNorm(config.hidden_size, - config.layernorm_epsilon, - quant_config=quantization_config, - dtype=dtype, - device=device, - prefix=add_prefix('post_attention_layernorm', prefix)) - self.mlp = ConceptLMV22VQOlmoMLP(config, dtype=dtype, device=device, prefix=add_prefix('mlp', prefix)) - self.post_feedforward_layernorm = RMSNorm(config.hidden_size, - config.layernorm_epsilon, - quant_config=quantization_config, - dtype=dtype, - device=device, - prefix=add_prefix('post_feedforward_layernorm', prefix)) - - def forward(self, - hidden_states: torch.Tensor, - rotary_pos_emb: tuple[torch.Tensor, torch.Tensor], - past_key_value: list[torch.Tensor] | None = None, - attn_metadata: Any = None): - """forward. - - The reference uses post-norm residuals:: - h = h + post_attention_layernorm(attn(h)) - h = h + post_feedforward_layernorm(mlp(h)) - This is NOT the same as lmdeploy's pre-norm residual form - ``norm(x, residual=h)`` (which normalizes x+h). So we call the norm - without a residual and add explicitly. - """ - attn_out = self.self_attention( - hidden_states=hidden_states, - rotary_pos_emb=rotary_pos_emb, - past_key_value=past_key_value, - attn_metadata=attn_metadata, - ) - hidden_states = hidden_states + self.post_attention_layernorm(attn_out) - - mlp_out = self.mlp(hidden_states) - hidden_states = hidden_states + self.post_feedforward_layernorm(mlp_out) - return hidden_states - - -class ConceptLMV22VQOlmoBlock(nn.Module): - """Rewrite of ``_OlmoBlock``. - - Holds a stack of ``_OlmoLayer`` and an optional final RMSNorm. Mirrors the - reference's ``forward`` return contract: ``(hidden_states, layer_states)``. - """ - - def __init__(self, - config: PretrainedConfig, - num_layers: int, - post_layer_norm: bool, - dtype: torch.dtype = None, - device: torch.device = None, - window_size: Any = _CONFIG_VALUE, - skip_frequency: Any = _CONFIG_VALUE, - prefix: str = ''): - super().__init__() - quantization_config = getattr(config, 'quantization_config', None) - if window_size is _CONFIG_VALUE: - window_size = _get_configured_window(config) - elif isinstance(window_size, (list, tuple)): - window_size = window_size[0] if len(window_size) > 0 else None - if window_size is not None: - window_size = int(window_size) - if skip_frequency is _CONFIG_VALUE: - skip_frequency = getattr(config, 'window_attn_skip_freq', None) - if skip_frequency is not None: - skip_frequency = int(skip_frequency) - - self.layers = nn.ModuleList([ - ConceptLMV22VQOlmoLayer(config, - layer_idx, - dtype=dtype, - device=device, - sliding_window=self._layer_sliding_window(layer_idx + 1, window_size, - skip_frequency), - prefix=add_prefix(f'layers.{layer_idx}', prefix)) - for layer_idx in range(num_layers) - ]) - self.final_layernorm = ( - RMSNorm(config.hidden_size, - config.layernorm_epsilon, - quant_config=quantization_config, - dtype=dtype, - device=device, - prefix=add_prefix('final_layernorm', prefix)) if post_layer_norm else None) - self.rotary_emb = _make_olmo_rotary_embedding(config, device=device) - self.num_heads = int(config.num_attention_heads) - self.head_dim = int(config.kv_channels) - - @staticmethod - def _layer_sliding_window(layer_number: int, window_size: int | None, skip_frequency: int | None): - """Match reference OLMo window/full attention alternation.""" - if window_size is None: - return None - if skip_frequency is not None and layer_number % skip_frequency == 0: - return None - return window_size - - def _make_rotary_pos_emb(self, hidden_states: torch.Tensor, position_ids: torch.Tensor): - """Create RoPE cos/sin in the shape consumed by ``ApplyRotaryEmb``.""" - if position_ids.dim() == 1: - position_ids = position_ids.unsqueeze(0) - cos, sin = self.rotary_emb(hidden_states, position_ids) - if cos.dim() == 3 and cos.size(0) == 1: - cos = cos[0] - sin = sin[0] - return cos, sin - - def forward(self, - hidden_states: torch.Tensor, - position_ids: torch.Tensor, - past_key_values: list[list[torch.Tensor]] | None = None, - attn_metadata: Any = None, - collect: bool = False): - """forward.""" - layer_states = [] - rotary_pos_emb = self._make_rotary_pos_emb(hidden_states, position_ids) - for idx, layer in enumerate(self.layers): - pkv = past_key_values[idx] if past_key_values is not None else None - hidden_states = layer( - hidden_states, - rotary_pos_emb=rotary_pos_emb, - past_key_value=pkv, - attn_metadata=attn_metadata, - ) - if collect: - layer_states.append(hidden_states) - if self.final_layernorm is not None: - hidden_states = self.final_layernorm(hidden_states) - return hidden_states, layer_states - - def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]], prefix: str = ''): - """Load native ConceptLM OLMo block weights into the LMDeploy - rewrite.""" - if prefix and not prefix.endswith('.'): - prefix = f'{prefix}.' - - params_dict = dict(self.named_parameters()) - loaded_names = set() - for name, loaded_weight in weights: - if 'rotary_emb.inv_freq' in name: - continue - if prefix: - if not name.startswith(prefix): - continue - name = name[len(prefix):] - - if name.endswith('.self_attention.qkv_proj.weight'): - param = params_dict[name] - query, key, value = param.weight_spliter(loaded_weight) - load_weight(param, query, shard_id='q') - load_weight(param, key, shard_id='k') - load_weight(param, value, shard_id='v') - loaded_names.add(name) - continue - - if name.endswith('.self_attention.linear_qkv.weight'): - target_name = name.replace('.self_attention.linear_qkv.weight', - '.self_attention.qkv_proj.weight') - param = params_dict[target_name] - loaded_weight = _repack_olmo_qkv_weight(loaded_weight, self.num_heads, self.head_dim) - query, key, value = param.weight_spliter(loaded_weight) - load_weight(param, query, shard_id='q') - load_weight(param, key, shard_id='k') - load_weight(param, value, shard_id='v') - loaded_names.add(target_name) - continue - - if name.endswith('.self_attention.linear_proj.weight'): - target_name = name.replace('.self_attention.linear_proj.weight', - '.self_attention.o_proj.weight') - elif name.endswith('.mlp.gate_up_proj.weight'): - param = params_dict[name] - gate, up = param.weight_spliter(loaded_weight) - load_weight(param, gate, shard_id=0) - load_weight(param, up, shard_id=1) - loaded_names.add(name) - continue - elif name.endswith('.mlp.linear_fc1.weight'): - target_name = name.replace('.mlp.linear_fc1.weight', '.mlp.gate_up_proj.weight') - param = params_dict[target_name] - gate, up = param.weight_spliter(loaded_weight) - load_weight(param, gate, shard_id=0) - load_weight(param, up, shard_id=1) - loaded_names.add(target_name) - continue - elif name.endswith('.mlp.linear_fc2.weight'): - target_name = name.replace('.mlp.linear_fc2.weight', '.mlp.down_proj.weight') - else: - target_name = name - - param = params_dict.get(target_name) - if param is None: - continue - load_weight(param, loaded_weight) - loaded_names.add(target_name) - return loaded_names - - -class ConceptLMV22VQForCausalLM(nn.Module, DeployModelMixinV1, CudaGraphMixin): - """Rewrote model of ConceptLMV22VQForCausalLM.""" - - def __init__(self, - config: PretrainedConfig, - ctx_mgr: StepContextManager, - dtype: torch.dtype = None, - device: torch.device = None, - prefix: str = ''): - super().__init__() - self.config = config - self.ctx_mgr = ctx_mgr - # token embedding — mirrors ``self.embedding`` in the reference. - self.embedding = ConceptLMV22VQEmbedding(config, dtype=dtype, device=device) - self.encoder = ConceptLMV22VQOlmoBlock(config, - config.concept_encoder_layers, - post_layer_norm=False, - dtype=dtype, - device=device, - prefix=add_prefix('encoder', prefix)) - self.decoder = ConceptLMV22VQOlmoBlock(config, - config.concept_decoder_layers, - post_layer_norm=True, - dtype=dtype, - device=device, - prefix=add_prefix('decoder', prefix)) - self.concept_vq_input_norm = nn.LayerNorm(config.hidden_size, - eps=getattr(config, 'layernorm_epsilon', 1e-6), - dtype=dtype, - device=device) - self.concept_quantizer = ConceptLMV22VQQuantizer(config, dtype=dtype, device=device) - self.concept_predictor = ConceptLMV22VQConceptPredictor(config, - dtype=dtype, - device=device, - prefix=add_prefix('concept_predictor', prefix)) - self.concept_predictor.set_attention_window(tuple(getattr(config, 'window_size', (None, None))), - getattr(config, 'window_attn_skip_freq', None)) - self.fusion_tok_norm = nn.LayerNorm(config.hidden_size, - eps=getattr(config, 'layernorm_epsilon', 1e-6), - dtype=dtype, - device=device) - self.fusion_hl_norm = nn.LayerNorm(config.hidden_size, - eps=getattr(config, 'layernorm_epsilon', 1e-6), - dtype=dtype, - device=device) - self.fusion_norm_alpha = nn.Parameter( - torch.tensor(getattr(config, 'concept_fusion_norm_alpha_init', 0.1), dtype=dtype, device=device), - requires_grad=False) - self.dd_encoder_self_dd = ConceptLMV22VQSelfDD(config, - config.concept_encoder_layers, - use_softmax=False, - dtype=dtype, - device=device) - self.decoder_read_encoder_routes = nn.ModuleList([ - ConceptLMV22VQResidualRoute(config, - config.concept_encoder_layers, - use_softmax=True, - dtype=dtype, - device=device) - for _ in range(config.concept_decoder_layers) - ]) - self.decoder_read_encoder_shared_source_norm = nn.LayerNorm(config.hidden_size, - eps=getattr(config, 'layernorm_epsilon', 1e-6), - dtype=dtype, - device=device) - self.decoder_read_concept_routes = nn.ModuleList([ - ConceptLMV22VQResidualRoute(config, - config.concept_special_layers, - use_softmax=True, - dtype=dtype, - device=device) - for _ in range(config.concept_decoder_layers) - ]) - self.decoder_read_concept_shared_source_norm = nn.LayerNorm(config.hidden_size, - eps=getattr(config, 'layernorm_epsilon', 1e-6), - dtype=dtype, - device=device) - self.final_read_concept_gate_logits = nn.Parameter( - torch.zeros(config.concept_decoder_layers, 2, dtype=dtype, device=device), - requires_grad=False) - self.dd_two_route_add = ConceptLMV22VQTwoRouteAdd(config, dtype=dtype, device=device) - self.concept_ops = ConceptLMRuntimeOps() - # output projection — mirrors ``output_layer`` in the reference. Built - # via build_lm_head and named ``lm_head`` so DeployModelMixinV1. - # get_logits picks it up directly; load_weights maps the checkpoint's - # ``output_layer`` onto it. - self.lm_head = self.build_lm_head( - config.hidden_size, config.vocab_size, bias=False, dtype=dtype, device=device) - - def forward(self, - input_ids: torch.Tensor, - position_ids: torch.Tensor, - past_key_values: list[list[torch.Tensor]], - attn_metadata=None, - inputs_embeds: torch.Tensor = None, - state_ids: torch.Tensor | None = None, - state_caches: list[torch.Tensor] | None = None, - **kwargs): - """Model forward, return hidden_states (logits computed by runtime).""" - concept_metadata = self._build_concept_metadata(position_ids, attn_metadata, state_ids) - concept_caches = self._build_concept_caches(past_key_values, state_caches) - if inputs_embeds is None: - hidden_states = self.embedding(input_ids) - else: - hidden_states = inputs_embeds - - if concept_metadata.is_decoding: - return self._forward_decode( - hidden_states, - position_ids, - concept_metadata, - concept_caches, - ) - - hidden_states, prefill_position_ids = self._normalize_prefill_inputs( - hidden_states, - position_ids, - attn_metadata, - ) - return self._forward_prefill_packed( - hidden_states, - prefill_position_ids, - concept_metadata, - concept_caches, - ) - - def get_input_embeddings(self): - """Get input embeddings.""" - return self.embedding.word_embeddings - - def get_output_embeddings(self): - """Get output embeddings.""" - return self.lm_head - - def _split_past_key_values(self, past_key_values: list[list[torch.Tensor]] | None): - """Split the flat LMDeploy KV-cache list into ConceptLM streams.""" - return _split_concept_past_key_values(self.config, past_key_values) - - def _build_concept_metadata(self, - position_ids: torch.Tensor | None, - attn_metadata: Any = None, - state_ids: torch.Tensor | None = None) -> ConceptMetadata: - """Build top-level ConceptLM metadata for future serving paths.""" - return ConceptMetadata.build( - self.config, - position_ids=position_ids, - attn_metadata=attn_metadata, - state_ids=state_ids, - ) - - def _build_concept_caches(self, - past_key_values: list[list[torch.Tensor]] | None, - state_caches: list[torch.Tensor] | None = None) -> ConceptCaches: - """Build top-level ConceptLM cache views for future serving paths.""" - return ConceptCaches.build( - self.config, - past_key_values=past_key_values, - state_caches=state_caches, - ) - - def _decode_chunk_state_update(self, - current_source_states: torch.Tensor, - concept_metadata: ConceptMetadata, - concept_caches: ConceptCaches) -> ConceptChunkStateUpdateResult: - """Update decode chunk-source state and return fixed-shape rows. - - CUDA uses the Triton writer. CPU uses the reference writer for tests. The returned rows deliberately avoid - dynamic concept-row compaction, matching the CUDA graph route in the design doc. - """ - chunk_source_state = concept_caches.chunk_source_state - if chunk_source_state is None: - raise RuntimeError('ConceptLM decode chunk update requires concept chunk source state cache.') - if concept_metadata.state_ids is None: - raise RuntimeError('ConceptLM decode chunk update requires state_ids.') - if concept_metadata.position_ids is None: - raise RuntimeError('ConceptLM decode chunk update requires position_ids.') - concept_input_states, next_rows, update_mask = self.concept_ops.decode_chunk_state_update( - chunk_source_state, - current_source_states, - concept_metadata.state_ids, - concept_metadata.position_ids, - concept_metadata.chunk_size, - concept_metadata.merge_method, - ) - state_ids = concept_metadata.state_ids.to(device=current_source_states.device, dtype=torch.long) - return ConceptChunkStateUpdateResult( - concept_input_states=concept_input_states, - next_chunk_source_states=next_rows, - concept_update_mask=update_mask, - valid_state_mask=state_ids >= 0, - state_ids=state_ids, - safe_state_ids=state_ids.clamp(min=0), - ) - - def support_cuda_graph( - self, - input_ids: torch.Tensor, - position_ids: torch.Tensor, - past_key_values: list[list[torch.Tensor]], - attn_metadata: Any = None, - inputs_embeds: torch.Tensor = None, - **kwargs, - ): - """Enable CUDA graph for fixed-shape decode only. - - Prefill still builds packed concept metadata with runtime compact sizes - and must stay eager. Decode boundary updates use fixed batch-shape - backend ops, so the CUDA graph runner can pad ``state_ids`` through the - common SSM path and replay the model with static tensor shapes. - """ - return bool(getattr(attn_metadata, 'is_decoding', False)) - - def _route_gate(self, layer_idx: int) -> torch.Tensor: - """Return decoder route gate ``[decoder_dd_scale, - concept_route_scale]``.""" - return self.final_read_concept_gate_logits[int(layer_idx)].float().softmax(dim=-1) - - def _normalize_prefill_inputs(self, - hidden_states: torch.Tensor, - position_ids: torch.Tensor, - attn_metadata: Any = None): - """Normalize LMDeploy prefill inputs to packed token layout. - - Engine layout stays fixed as ``[1, total_tokens, hidden]`` with - ``position_ids=[1, total_tokens]``. Per-request boundaries come from - ``attn_metadata.q_seqlens``/``q_start_loc`` and are used only to build - the concept stream. - """ - if attn_metadata is None or getattr(attn_metadata, 'q_seqlens', None) is None: - raise RuntimeError('ConceptLM prefill requires q_seqlens attention metadata.') - if getattr(attn_metadata, 'is_decoding', False): - raise RuntimeError('ConceptLM prefill input normalization received decode metadata.') - - if hidden_states.dim() != 3 or hidden_states.size(0) != 1: - raise NotImplementedError( - f'ConceptLM prefill expects fixed engine layout [1, total_tokens, hidden], ' - f'got {tuple(hidden_states.shape)}.') - - total_tokens = hidden_states.size(1) - if position_ids.dim() != 2 or position_ids.size(0) != 1: - raise NotImplementedError( - f'ConceptLM prefill expects fixed position layout [1, total_tokens], ' - f'got {tuple(position_ids.shape)}.') - if position_ids.size(1) != total_tokens: - raise ValueError(f'position_ids length {position_ids.size(1)} does not match token length {total_tokens}.') - - q_seqlens = attn_metadata.q_seqlens - if q_seqlens.dim() != 1: - raise ValueError(f'ConceptLM prefill expects 1-D q_seqlens, got {tuple(q_seqlens.shape)}.') - - hidden_states = hidden_states[0].contiguous() - return hidden_states, position_ids[0].to(device=hidden_states.device) - - def _normalize_decode_inputs(self, hidden_states: torch.Tensor, position_ids: torch.Tensor): - """Normalize LMDeploy decode inputs to one flat row per active - request.""" - if hidden_states.dim() != 3 or hidden_states.size(0) != 1: - raise NotImplementedError( - f'ConceptLM decode expects fixed engine layout [1, batch, hidden], ' - f'got {tuple(hidden_states.shape)}.') - batch_size = hidden_states.size(1) - position_ids = _flatten_decode_position_ids(position_ids, batch_size).to(device=hidden_states.device) - if position_ids.numel() != batch_size: - raise ValueError(f'Expected {batch_size} decode position ids, got {position_ids.numel()}.') - return hidden_states[0].contiguous(), position_ids - - @staticmethod - def _build_decode_metadata(position_ids: torch.Tensor, - state_ids: torch.Tensor | None, - batch_size: int, - device: torch.device) -> ConceptDecodeMetadata: - """Build fixed-shape decode metadata from engine state ids.""" - if state_ids is None: - raise RuntimeError('ConceptLM decode requires state_ids.') - state_ids = state_ids.to(device=device, dtype=torch.long).reshape(-1) - if state_ids.numel() != batch_size: - raise ValueError(f'Expected {batch_size} decode state ids, got {state_ids.numel()}.') - valid_state_mask = state_ids >= 0 - return ConceptDecodeMetadata( - position_ids=position_ids, - state_ids=state_ids, - safe_state_ids=state_ids.clamp(min=0), - valid_state_mask=valid_state_mask, - ) - - @staticmethod - def _select_decode_state_rows(state_cache: torch.Tensor, - decode_metadata: ConceptDecodeMetadata) -> torch.Tensor: - """Gather state-cache rows and zero out padded decode rows.""" - rows = state_cache.index_select(0, decode_metadata.safe_state_ids) - mask_shape = (decode_metadata.valid_state_mask.size(0), ) + (1, ) * (rows.dim() - 1) - valid_mask = decode_metadata.valid_state_mask.view(mask_shape) - return torch.where(valid_mask, rows, torch.zeros_like(rows)) - - def _decode_concept_read_mask(self, decode_metadata: ConceptDecodeMetadata) -> torch.Tensor: - """Return rows whose current decode token should read a cached - concept.""" - repeat_slots = self._repeat_slot_ids( - decode_metadata.position_ids, - int(self.config.concept_chunk_size), - bool(getattr(self.config, 'concept_shift_feature', True)), - ) - return decode_metadata.valid_state_mask & (repeat_slots >= 0) - - def _build_decode_chunk_source_states(self, - hidden_states: torch.Tensor, - encoder_raw_states: list[torch.Tensor]) -> torch.Tensor: - """Build current per-row states accumulated until the next concept - boundary. - - Row 0 is the final encoder hidden, used as concept-predictor input when - a boundary is reached. Remaining rows mirror the prefill - ``encoder_raw_states[:-1]`` route sources. - """ - source_states = [hidden_states] - source_states.extend(encoder_raw_states[:-1]) - return torch.stack(tuple(source_states), dim=1) - - @staticmethod - def _decode_concept_position_ids(position_ids: torch.Tensor, chunk_size: int) -> torch.Tensor: - """Return reference RoPE positions for concept rows emitted at decode - boundaries.""" - return (position_ids - int(chunk_size) + 1).clamp(min=0) - - def _build_concept_decode_metadata_static(self, - token_attn_metadata: Any, - decode_metadata: ConceptDecodeMetadata): - """Build fixed-shape concept-stream decode metadata. - - Concept predictor runs with the same batch shape as token decode. - Boundary rows append a real concept KV entry. Non-boundary/padded rows - execute dummy concept attention with safe ``kv_seqlens >= 1``; their KV - writes are restored afterward by ``ConceptLMRuntimeOps``. - """ - if token_attn_metadata is None or getattr(token_attn_metadata, 'block_offsets', None) is None: - raise RuntimeError('ConceptLM decode concept update requires token attention metadata.') - - device = decode_metadata.position_ids.device - batch_size = decode_metadata.position_ids.numel() - q_seqlens = getattr(token_attn_metadata, 'q_seqlens', None) - q_start_loc = getattr(token_attn_metadata, 'q_start_loc', None) - kv_seqlens = getattr(token_attn_metadata, 'kv_seqlens', None) - if q_seqlens is None or q_start_loc is None or kv_seqlens is None: - raise RuntimeError('ConceptLM decode concept update requires q/q_start/kv sequence metadata.') - q_dtype = q_seqlens.dtype - q_start_dtype = q_start_loc.dtype - kv_dtype = kv_seqlens.dtype - - concept_q_seqlens = torch.ones((batch_size, ), dtype=q_dtype, device=device) - concept_q_start_loc = torch.arange(batch_size, dtype=q_start_dtype, device=device) - concept_cu_seqlens = F.pad(torch.cumsum(concept_q_seqlens, dim=0, dtype=torch.int32), (1, 0)) - concept_kv_seqlens = torch.div( - decode_metadata.position_ids + 1, - int(self.config.concept_chunk_size), - rounding_mode='floor', - ).clamp(min=1).to(dtype=kv_dtype) - - updates = dict( - is_decoding=True, - block_offsets=token_attn_metadata.block_offsets, - q_start_loc=concept_q_start_loc, - q_seqlens=concept_q_seqlens, - kv_seqlens=concept_kv_seqlens, - cu_seqlens_q=concept_cu_seqlens, - cu_seqlens_k=concept_cu_seqlens, - ) - if hasattr(token_attn_metadata, 'kv_start_loc'): - updates['kv_start_loc'] = concept_kv_seqlens - concept_q_seqlens.to(dtype=concept_kv_seqlens.dtype) - if hasattr(token_attn_metadata, 'kv_flatten_size'): - updates['kv_flatten_size'] = batch_size - if hasattr(token_attn_metadata, 'max_q_seqlen'): - updates['max_q_seqlen'] = 1 - if hasattr(token_attn_metadata, 'max_kv_seqlen'): - # Decode kernels use kv_seqlens directly. Keep a conservative bound - # without reading the dynamic maximum back to host. - updates['max_kv_seqlen'] = getattr(token_attn_metadata, 'max_kv_seqlen') - if hasattr(token_attn_metadata, 'scheduler_metadata'): - updates['scheduler_metadata'] = None - if hasattr(token_attn_metadata, 'tile_scheduler_metadata'): - updates['tile_scheduler_metadata'] = None - if hasattr(token_attn_metadata, 'num_splits'): - updates['num_splits'] = None - if hasattr(token_attn_metadata, 'fill_seqlens'): - updates['fill_seqlens'] = concept_q_seqlens - return replace(token_attn_metadata, **updates) - - @staticmethod - def _stack_concept_raw_states(concept_raw_states: list[torch.Tensor]) -> torch.Tensor: - """Stack raw concept-layer states to ``[rows, concept_layers, - hidden]``.""" - return torch.stack(tuple(concept_raw_states), dim=1) - - def _snapshot_decode_concept_kv(self, - concept_caches: ConceptCaches, - concept_attn_metadata: Any): - """Snapshot concept KV slots that dummy non-boundary rows may - overwrite.""" - if concept_caches.concept_past_key_values is None: - raise RuntimeError('ConceptLM decode concept update requires concept KV caches.') - return [ - self.concept_ops.decode_kv_cache_snapshot( - k_cache, - v_cache, - concept_attn_metadata.block_offsets, - concept_attn_metadata.kv_seqlens, - ) - for k_cache, v_cache in concept_caches.concept_past_key_values - ] - - def _restore_decode_concept_kv_( - self, - concept_caches: ConceptCaches, - concept_attn_metadata: Any, - saved_kv, - restore_mask: torch.Tensor, - ): - """Restore concept KV slots for non-boundary and padded rows.""" - if concept_caches.concept_past_key_values is None: - raise RuntimeError('ConceptLM decode concept update requires concept KV caches.') - for (k_cache, v_cache), (saved_k, saved_v) in zip(concept_caches.concept_past_key_values, saved_kv): - self.concept_ops.decode_kv_cache_restore( - k_cache, - v_cache, - saved_k, - saved_v, - concept_attn_metadata.block_offsets, - concept_attn_metadata.kv_seqlens, - restore_mask, - ) - - def _write_decode_concept_states_static_(self, - concept_caches: ConceptCaches, - decode_metadata: ConceptDecodeMetadata, - update_mask: torch.Tensor, - predicted_vectors: torch.Tensor, - concept_raw_states: list[torch.Tensor]): - """Write newly emitted concept states to persistent decode caches.""" - last_final_state = concept_caches.last_final_state - last_raw_states = concept_caches.last_raw_states - if last_final_state is None or last_raw_states is None: - raise RuntimeError('ConceptLM decode concept update requires last concept state caches.') - raw_rows = self._stack_concept_raw_states(concept_raw_states) - self.concept_ops.decode_concept_state_update( - last_raw_states, - last_final_state, - predicted_vectors, - raw_rows, - decode_metadata.state_ids, - update_mask, - ) - - def _update_decode_concept_states_static_(self, - chunk_update: ConceptChunkStateUpdateResult, - decode_metadata: ConceptDecodeMetadata, - concept_metadata: ConceptMetadata, - concept_caches: ConceptCaches): - """Emit/cache concept states with fixed batch shape. - - The predictor runs for every decode row so CUDA graph capture sees a stable launch sequence. Non-boundary rows - are dummy work: their concept KV writes are restored and their final/raw state writes are masked. - """ - if concept_caches.concept_past_key_values is None: - raise RuntimeError('ConceptLM decode concept update requires concept KV caches.') - - concept_hidden = self.concept_vq_input_norm(chunk_update.concept_input_states[:, 0]) - encoder_concept_states = self.concept_predictor.normalize_encoder_concept_states( - chunk_update.concept_input_states[:, 1:]) - concept_position_ids = self._decode_concept_position_ids( - decode_metadata.position_ids, - concept_metadata.chunk_size, - ) - concept_attn_metadata = self._build_concept_decode_metadata_static( - concept_metadata.attn_metadata, - decode_metadata, - ) - saved_kv = self._snapshot_decode_concept_kv(concept_caches, concept_attn_metadata) - concept_logits, concept_raw_states = self.concept_predictor( - concept_hidden, - encoder_concept_states, - concept_position_ids, - past_key_values=concept_caches.concept_past_key_values, - attn_metadata=concept_attn_metadata, - ) - self._restore_decode_concept_kv_( - concept_caches, - concept_attn_metadata, - saved_kv, - ~chunk_update.concept_update_mask, - ) - predicted_vectors = self.concept_quantizer(concept_logits) - self._write_decode_concept_states_static_( - concept_caches, - decode_metadata, - chunk_update.concept_update_mask, - predicted_vectors, - concept_raw_states, - ) - - def _merge_prefill_tail_chunk_states(self, - source_states: torch.Tensor, - prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: - """Build per-request partial chunk accumulator rows after prefill.""" - assert source_states.dim() == 3, ( - f'Expected source states [total_tokens, num_sources, hidden], got {tuple(source_states.shape)}.') - device = source_states.device - chunk_size = int(self.config.concept_chunk_size) - q_seqlens = prefill_metadata.token_q_seqlens.to(device=device, dtype=torch.long) - q_start_loc = prefill_metadata.token_q_start_loc.to(device=device, dtype=torch.long) - batch_size = q_seqlens.size(0) - tail_lens = torch.remainder(q_seqlens, chunk_size) - tail_lens = torch.where(q_seqlens < chunk_size, q_seqlens, tail_lens) - tail_lens = torch.where(q_seqlens > 0, tail_lens, torch.zeros_like(tail_lens)) - has_tail = tail_lens > 0 - - tail_rows = source_states.new_zeros((batch_size, source_states.size(1), source_states.size(2))) - if source_states.size(0) == 0: - return tail_rows - merge_method = getattr(self.config, 'concept_chunk_merge_method', 'meanpooling') - if merge_method == 'first': - token_ids = q_start_loc + q_seqlens - tail_lens - token_ids = token_ids.clamp(min=0, max=max(source_states.size(0) - 1, 0)) - rows = source_states[token_ids] - return torch.where(has_tail.view(batch_size, 1, 1), rows, tail_rows) - if merge_method == 'last': - token_ids = q_start_loc + q_seqlens - 1 - token_ids = token_ids.clamp(min=0, max=max(source_states.size(0) - 1, 0)) - rows = source_states[token_ids] - return torch.where(has_tail.view(batch_size, 1, 1), rows, tail_rows) - - token_ids = torch.arange(prefill_metadata.num_tokens_total, dtype=torch.long, device=device) - cu_q_seqlens = F.pad(torch.cumsum(q_seqlens, dim=0), (1, 0)) - token_seq = torch.searchsorted(cu_q_seqlens[1:], token_ids, right=True) - token_pos = token_ids - cu_q_seqlens[token_seq] - token_tail_start = q_seqlens[token_seq] - tail_lens[token_seq] - valid_tail = (tail_lens[token_seq] > 0) & (token_pos >= token_tail_start) - weighted_source = source_states * valid_tail.to(dtype=source_states.dtype).view(-1, 1, 1) - tail_rows.index_add_(0, token_seq, weighted_source) - return tail_rows - - def _write_prefill_state_caches_eager_(self, - concept_caches: ConceptCaches, - concept_metadata: ConceptMetadata, - prefill_metadata: ConceptPrefillMetadata, - source_states: torch.Tensor, - predicted_vectors: torch.Tensor, - concept_raw_states: list[torch.Tensor]): - """Seed decode state caches from a completed prefill forward.""" - if concept_caches.state_caches is None or concept_metadata.state_ids is None: - return - chunk_source_state = concept_caches.chunk_source_state - last_raw_states = concept_caches.last_raw_states - last_final_state = concept_caches.last_final_state - if chunk_source_state is None or last_raw_states is None or last_final_state is None: - raise RuntimeError('ConceptLM prefill state init requires all ConceptLM state caches.') - - state_ids = concept_metadata.state_ids.to(device=source_states.device, dtype=torch.long).reshape(-1) - batch_size = prefill_metadata.token_q_seqlens.numel() - if state_ids.numel() != batch_size: - raise ValueError(f'Expected {batch_size} prefill state ids, got {state_ids.numel()}.') - valid_indices = torch.nonzero(state_ids >= 0, as_tuple=False).flatten() - if valid_indices.numel() == 0: - return - - valid_state_ids = state_ids.index_select(0, valid_indices) - tail_rows = self._merge_prefill_tail_chunk_states(source_states, prefill_metadata) - chunk_source_state.index_copy_(0, valid_state_ids, - tail_rows.index_select(0, valid_indices).to(dtype=chunk_source_state.dtype)) - - concept_counts = prefill_metadata.concept_q_seqlens.to(device=source_states.device, dtype=torch.long) - concept_start = prefill_metadata.concept_q_start_loc.to(device=source_states.device, dtype=torch.long) - concept_valid_mask = (state_ids >= 0) & (concept_counts > 0) - concept_indices = torch.nonzero(concept_valid_mask, as_tuple=False).flatten() - if concept_indices.numel() == 0: - return - - concept_state_ids = state_ids.index_select(0, concept_indices) - last_concept_ids = concept_start.index_select(0, concept_indices) + concept_counts.index_select( - 0, concept_indices) - 1 - last_final_rows = predicted_vectors.index_select(0, last_concept_ids) - last_final_state.index_copy_(0, concept_state_ids, last_final_rows.to(dtype=last_final_state.dtype)) - raw_rows = self._stack_concept_raw_states(concept_raw_states).index_select(0, last_concept_ids) - last_raw_states.index_copy_(0, concept_state_ids, raw_rows.to(dtype=last_raw_states.dtype)) - - @staticmethod - def _concept_count_from_seq_len(seq_len: int, chunk_size: int) -> int: - """Return reference ConceptLM chunk count for one request length.""" - seq_len = int(seq_len or 0) - if seq_len <= 0: - return 0 - if seq_len < chunk_size: - return 1 - return seq_len // chunk_size - - @staticmethod - def _concept_counts_from_q_seqlens(q_seqlens: torch.Tensor, chunk_size: int) -> torch.Tensor: - """Vectorized version of ``_concept_count_from_seq_len``.""" - counts = torch.div(q_seqlens, chunk_size, rounding_mode='floor').clamp(min=1) - return torch.where(q_seqlens > 0, counts, torch.zeros_like(q_seqlens)) - - @staticmethod - def _repeat_slot_ids(token_pos: torch.Tensor, chunk_size: int, shift_feature: bool) -> torch.Tensor: - """Return local concept slot read by each token after shift - semantics.""" - if shift_feature: - return torch.div(token_pos + 1, chunk_size, rounding_mode='floor') - 1 - return torch.div(token_pos, chunk_size, rounding_mode='floor') - 1 - - @staticmethod - def _get_max_concepts_per_request(token_attn_metadata: Any, - concept_q_seqlens: torch.Tensor, - chunk_size: int) -> int: - """Return per-request concept attention bound without hidden context - access.""" - max_q_seqlen = getattr(token_attn_metadata, 'max_q_seqlen', None) - if max_q_seqlen is not None: - return ConceptLMV22VQForCausalLM._concept_count_from_seq_len(int(max_q_seqlen), chunk_size) - - # Test/direct-call fallback. Serving should use the scheduler-provided - # Python max_q_seqlen above, as in DSV4 metadata construction. - return int(concept_q_seqlens.max().item()) if concept_q_seqlens.numel() > 0 else 0 - - def _build_prefill_metadata(self, - token_attn_metadata: Any, - position_ids: torch.Tensor) -> ConceptPrefillMetadata: - """Build packed token-to-concept metadata for batched prefill.""" - if token_attn_metadata is None: - raise RuntimeError('ConceptLM prefill requires attention metadata.') - if getattr(token_attn_metadata, 'is_decoding', False): - raise RuntimeError('ConceptLM prefill metadata cannot be built from decode metadata.') - - q_seqlens = token_attn_metadata.q_seqlens - q_start_loc = getattr(token_attn_metadata, 'q_start_loc', None) - if q_start_loc is None: - q_start_loc = F.pad(torch.cumsum(q_seqlens, dim=0, dtype=torch.int32), (1, 0))[:-1] - - total_tokens = int(position_ids.numel()) - chunk_size = int(self.config.concept_chunk_size) - shift_feature = bool(getattr(self.config, 'concept_shift_feature', True)) - - q_seqlens_long = q_seqlens.to(dtype=torch.long, device=position_ids.device) - q_start_loc_long = q_start_loc.to(dtype=torch.long, device=position_ids.device) - cu_q_seqlens = getattr(token_attn_metadata, 'cu_seqlens_q', None) - if cu_q_seqlens is None: - cu_q_seqlens = F.pad(torch.cumsum(q_seqlens, dim=0, dtype=torch.int32), (1, 0)) - cu_q_seqlens_long = cu_q_seqlens.to(dtype=torch.long, device=position_ids.device) - - token_ids = torch.arange(total_tokens, dtype=torch.long, device=position_ids.device) - token_seq = torch.searchsorted(cu_q_seqlens_long[1:], token_ids, right=True) - token_pos = token_ids - cu_q_seqlens_long[token_seq] - - concept_q_seqlens_long = self._concept_counts_from_q_seqlens(q_seqlens_long, chunk_size) - concept_q_seqlens = concept_q_seqlens_long.to(dtype=q_seqlens.dtype, device=q_seqlens.device) - concept_q_start_loc = F.pad(torch.cumsum(concept_q_seqlens, dim=0, dtype=torch.int32), (1, 0))[:-1] - concept_q_start_loc_long = concept_q_start_loc.to(dtype=torch.long, device=position_ids.device) - concept_cu_seqlens_long = F.pad(torch.cumsum(concept_q_seqlens_long, dim=0), (1, 0)) - - # TODO: remove this eager compact-size scalar read by moving ConceptLM - # concept-stream allocation/compaction into the engine/backend contract, - # like DSV4's precomputed metadata and Qwen3.5's state-cache metadata. - num_concepts_total = int(concept_cu_seqlens_long[-1].item()) - concept_ids = torch.arange(num_concepts_total, dtype=torch.long, device=position_ids.device) - concept_seq = torch.searchsorted(concept_cu_seqlens_long[1:], concept_ids, right=True) - local_concept_ids = concept_ids - concept_cu_seqlens_long[concept_seq] - concept_token_start = q_start_loc_long[concept_seq] + local_concept_ids * chunk_size - concept_position_ids = position_ids[concept_token_start] - - seq_concept_start = concept_q_start_loc_long[token_seq] - seq_concept_count = concept_q_seqlens_long[token_seq] - repeat_slots = self._repeat_slot_ids(token_pos, chunk_size, shift_feature) - valid_repeat = (repeat_slots >= 0) & (repeat_slots < seq_concept_count) - token_to_concept = torch.where( - valid_repeat, - seq_concept_start + repeat_slots, - torch.full_like(repeat_slots, -1), - ) - - merge_slots = torch.div(token_pos, chunk_size, rounding_mode='floor') - valid_merge = (merge_slots >= 0) & (merge_slots < seq_concept_count) - merge_token_to_concept = torch.where( - valid_merge, - seq_concept_start + merge_slots, - torch.full_like(merge_slots, -1), - ) - safe_merge_ids = merge_token_to_concept.clamp(min=0) - merge_token_counts = torch.zeros(num_concepts_total, dtype=torch.int32, device=position_ids.device) - merge_token_counts.index_add_(0, safe_merge_ids, valid_merge.to(dtype=torch.int32)) - - concept_seq_len = q_seqlens_long[concept_seq] - merge_first_pos = torch.where( - concept_seq_len < chunk_size, - torch.zeros_like(local_concept_ids), - local_concept_ids * chunk_size, - ) - merge_last_pos = torch.where( - concept_seq_len < chunk_size, - (concept_seq_len - 1).clamp(min=0), - local_concept_ids * chunk_size + chunk_size - 1, - ) - merge_first_token_ids = q_start_loc_long[concept_seq] + merge_first_pos - merge_last_token_ids = q_start_loc_long[concept_seq] + merge_last_pos - merge_short_concept_mask = concept_seq_len < chunk_size - max_concepts_per_request = self._get_max_concepts_per_request( - token_attn_metadata, - concept_q_seqlens_long, - chunk_size, - ) - - return ConceptPrefillMetadata( - token_q_seqlens=q_seqlens, - token_q_start_loc=q_start_loc, - concept_q_seqlens=concept_q_seqlens, - concept_q_start_loc=concept_q_start_loc, - concept_position_ids=concept_position_ids, - merge_token_to_concept=merge_token_to_concept, - merge_token_counts=merge_token_counts, - merge_first_token_ids=merge_first_token_ids, - merge_last_token_ids=merge_last_token_ids, - merge_short_concept_mask=merge_short_concept_mask, - token_to_concept=token_to_concept, - num_tokens_total=total_tokens, - num_concepts_total=num_concepts_total, - max_concepts_per_request=max_concepts_per_request, - ) - - @staticmethod - def _merge_chunks_mean_packed(hidden_states: torch.Tensor, - prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: - """Mean-pool packed token states by precomputed concept ids.""" - merge_token_to_concept = prefill_metadata.merge_token_to_concept.to(device=hidden_states.device) - valid_merge = (merge_token_to_concept >= 0).to(dtype=hidden_states.dtype).unsqueeze(-1) - safe_merge_ids = merge_token_to_concept.clamp(min=0) - merged = hidden_states.new_zeros((prefill_metadata.num_concepts_total, hidden_states.size(-1))) - merged.index_add_(0, safe_merge_ids, hidden_states * valid_merge) - counts = prefill_metadata.merge_token_counts.clamp(min=1).to(device=hidden_states.device, - dtype=hidden_states.dtype) - return merged / counts.unsqueeze(-1) - - def _merge_chunks_packed(self, - hidden_states: torch.Tensor, - prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: - """Merge packed token states into packed per-request concept states.""" - assert hidden_states.dim() == 2, ( - f'_merge_chunks_packed expects [total_tokens, hidden], got {tuple(hidden_states.shape)}.') - merge_method = getattr(self.config, 'concept_chunk_merge_method', 'meanpooling') - if prefill_metadata.num_concepts_total <= 0: - raise RuntimeError('ConceptLM prefill produced no concept chunks.') - if merge_method == 'first': - merged = hidden_states[prefill_metadata.merge_first_token_ids] - short_mean = self._merge_chunks_mean_packed(hidden_states, prefill_metadata) - return torch.where(prefill_metadata.merge_short_concept_mask[:, None], short_mean, merged) - if merge_method == 'last': - merged = hidden_states[prefill_metadata.merge_last_token_ids] - short_mean = self._merge_chunks_mean_packed(hidden_states, prefill_metadata) - return torch.where(prefill_metadata.merge_short_concept_mask[:, None], short_mean, merged) - - return self._merge_chunks_mean_packed(hidden_states, prefill_metadata) - - @staticmethod - def _gather_zero_prefixed_concepts(concept_states_with_zero: torch.Tensor, - prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: - """Gather zero-prefixed concept rows to packed token rows.""" - token_to_concept = prefill_metadata.token_to_concept.to(device=concept_states_with_zero.device) - gather_ids = torch.clamp(token_to_concept + 1, min=0) - return concept_states_with_zero[gather_ids] - - def _repeat_shift_packed(self, - concept_states: torch.Tensor, - prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: - """Gather packed concept states back to packed token states.""" - concept_states_with_zero = torch.cat((torch.zeros_like(concept_states[:1]), concept_states), dim=0) - return self._gather_zero_prefixed_concepts(concept_states_with_zero, prefill_metadata) - - def _repeat_shift_source_states_packed(self, - concept_states_with_zero: torch.Tensor, - prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: - """Gather zero-prefixed packed concept source states to token rows.""" - return self._gather_zero_prefixed_concepts(concept_states_with_zero, prefill_metadata) - - def _build_encoder_concept_states_packed(self, - encoder_raw_states: list[torch.Tensor], - prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: - """Build packed chunk-level encoder states used by the concept - predictor.""" - chunks = [self._merge_chunks_packed(state, prefill_metadata) for state in encoder_raw_states[:-1]] - assert len(chunks) > 0, 'ConceptLM concept-read-encoder route requires at least one encoder source state.' - states = torch.stack(chunks, dim=-2) - return self.concept_predictor.normalize_encoder_concept_states(states) - - def _build_concept_prefill_metadata(self, - token_attn_metadata: Any, - prefill_metadata: ConceptPrefillMetadata): - """Build chunk-stream attention metadata for packed prefill.""" - if token_attn_metadata is None: - raise RuntimeError('ConceptLM prefill requires attention metadata for concept predictor attention.') - if getattr(token_attn_metadata, 'is_decoding', False): - raise RuntimeError('ConceptLM concept prefill metadata cannot be built from decode metadata.') - - concept_q_seqlens = prefill_metadata.concept_q_seqlens - concept_q_start_loc = prefill_metadata.concept_q_start_loc - concept_cu_seqlens = torch.nn.functional.pad( - torch.cumsum(concept_q_seqlens, dim=0, dtype=torch.int32), - (1, 0), - ) - max_concept_seqlen = int(prefill_metadata.max_concepts_per_request) - - updates = dict( - is_decoding=False, - q_start_loc=concept_q_start_loc, - q_seqlens=concept_q_seqlens, - kv_seqlens=concept_q_seqlens, - cu_seqlens_q=concept_cu_seqlens, - cu_seqlens_k=concept_cu_seqlens, - ) - if hasattr(token_attn_metadata, 'kv_start_loc'): - updates['kv_start_loc'] = concept_q_start_loc - if hasattr(token_attn_metadata, 'kv_flatten_size'): - updates['kv_flatten_size'] = int(prefill_metadata.num_concepts_total) - if hasattr(token_attn_metadata, 'max_q_seqlen'): - updates['max_q_seqlen'] = max_concept_seqlen - if hasattr(token_attn_metadata, 'max_kv_seqlen'): - updates['max_kv_seqlen'] = max_concept_seqlen - return replace(token_attn_metadata, **updates) - - def _forward_prefill_packed(self, - hidden_states: torch.Tensor, - position_ids: torch.Tensor, - concept_metadata: ConceptMetadata, - concept_caches: ConceptCaches): - """Packed non-decode ConceptLM forward.""" - if (concept_caches.encoder_past_key_values is None or concept_caches.concept_past_key_values is None - or concept_caches.decoder_past_key_values is None): - raise RuntimeError('ConceptLM prefill requires encoder, concept, and decoder KV caches.') - prefill_metadata = self._build_prefill_metadata(concept_metadata.attn_metadata, position_ids) - hidden_states, encoder_raw_states = self._encode( - hidden_states, - position_ids, - past_key_values=concept_caches.encoder_past_key_values, - attn_metadata=concept_metadata.attn_metadata, - ) - concept_hidden = self._merge_chunks_packed(hidden_states, prefill_metadata) - concept_hidden = self.concept_vq_input_norm(concept_hidden) - encoder_concept_states = self._build_encoder_concept_states_packed(encoder_raw_states, prefill_metadata) - concept_attn_metadata = self._build_concept_prefill_metadata( - concept_metadata.attn_metadata, - prefill_metadata, - ) - concept_logits, concept_raw_states = self.concept_predictor( - concept_hidden, - encoder_concept_states, - prefill_metadata.concept_position_ids, - past_key_values=concept_caches.concept_past_key_values, - attn_metadata=concept_attn_metadata, - ) - predicted_vectors = self.concept_quantizer(concept_logits) - self._write_prefill_state_caches_eager_( - concept_caches, - concept_metadata, - prefill_metadata, - self._build_decode_chunk_source_states(hidden_states, encoder_raw_states), - predicted_vectors, - concept_raw_states, - ) - repeated_concepts = self._repeat_shift_packed(predicted_vectors, prefill_metadata) - decoder_input = self.fusion_tok_norm(hidden_states) + self.fusion_norm_alpha.to( - hidden_states.dtype) * self.fusion_hl_norm(repeated_concepts.to(hidden_states.dtype)) - final_hidden = self._decode( - decoder_input, - encoder_raw_states, - repeated_concepts, - concept_raw_states, - position_ids, - past_key_values=concept_caches.decoder_past_key_values, - attn_metadata=concept_metadata.attn_metadata, - prefill_metadata=prefill_metadata, - ) - return final_hidden.unsqueeze(0).contiguous() - - def _forward_decode(self, - hidden_states: torch.Tensor, - position_ids: torch.Tensor, - concept_metadata: ConceptMetadata, - concept_caches: ConceptCaches): - """ConceptLM decode path. - - This path is semantically structured for serving. Boundary concept - updates run with fixed batch shape so it is eligible for CUDA graph - replay through ``support_cuda_graph``. - """ - if concept_caches.encoder_past_key_values is None or concept_caches.decoder_past_key_values is None: - raise RuntimeError('ConceptLM decode requires encoder and decoder KV caches.') - if concept_caches.chunk_source_state is None: - raise RuntimeError('ConceptLM decode requires chunk source state cache.') - if concept_caches.last_raw_states is None or concept_caches.last_final_state is None: - raise RuntimeError('ConceptLM decode requires cached last concept states.') - - hidden_states, decode_position_ids = self._normalize_decode_inputs(hidden_states, position_ids) - decode_metadata = self._build_decode_metadata( - decode_position_ids, - concept_metadata.state_ids, - hidden_states.size(0), - hidden_states.device, - ) - decode_concept_metadata = replace( - concept_metadata, - position_ids=decode_position_ids, - state_ids=decode_metadata.state_ids, - ) - - hidden_states, encoder_raw_states = self._encode( - hidden_states, - decode_position_ids, - past_key_values=concept_caches.encoder_past_key_values, - attn_metadata=concept_metadata.attn_metadata, - ) - previous_final_concept_state = self._select_decode_state_rows(concept_caches.last_final_state, decode_metadata) - previous_concept_raw_state_rows = self._select_decode_state_rows(concept_caches.last_raw_states, - decode_metadata) - current_source_states = self._build_decode_chunk_source_states(hidden_states, encoder_raw_states) - chunk_update = self._decode_chunk_state_update( - current_source_states, - decode_concept_metadata, - concept_caches, - ) - self._update_decode_concept_states_static_( - chunk_update, - decode_metadata, - decode_concept_metadata, - concept_caches, - ) - - if bool(getattr(self.config, 'concept_shift_feature', True)): - final_concept_state = self._select_decode_state_rows(concept_caches.last_final_state, decode_metadata) - concept_raw_state_rows = self._select_decode_state_rows(concept_caches.last_raw_states, decode_metadata) - else: - final_concept_state = previous_final_concept_state - concept_raw_state_rows = previous_concept_raw_state_rows - concept_read_mask = self._decode_concept_read_mask(decode_metadata) - final_concept_state = torch.where(concept_read_mask.view(-1, 1), final_concept_state, - torch.zeros_like(final_concept_state)) - concept_raw_state_rows = torch.where(concept_read_mask.view(-1, 1, 1), concept_raw_state_rows, - torch.zeros_like(concept_raw_state_rows)) - decoder_input = self.fusion_tok_norm(hidden_states) + self.fusion_norm_alpha.to( - hidden_states.dtype) * self.fusion_hl_norm(final_concept_state.to(hidden_states.dtype)) - final_hidden = self._decode( - decoder_input, - encoder_raw_states, - final_concept_state, - concept_raw_states=[], - position_ids=decode_position_ids, - past_key_values=concept_caches.decoder_past_key_values, - attn_metadata=concept_metadata.attn_metadata, - decode_concept_states=concept_raw_state_rows, - ) - return final_hidden.unsqueeze(0).contiguous() - - def _encode(self, - hidden_states: torch.Tensor, - position_ids: torch.Tensor, - past_key_values: list[list[torch.Tensor]] | None = None, - attn_metadata: Any = None): - """Encoder stack plus encoder self-DD. - - This helper mirrors the reference flow but consumes LMDeploy attention inputs. It is wired for the future full - forward path; continuous batching still needs caller-side concept metadata before the full model can use it - safely. - """ - raw_states = [] - history_buffer = self.dd_encoder_self_dd.make_history_buffer(hidden_states) - self.dd_encoder_self_dd.write_history(history_buffer, 0, hidden_states) - rotary_pos_emb = self.encoder._make_rotary_pos_emb(hidden_states, position_ids) - - for layer_idx, layer in enumerate(self.encoder.layers): - pkv = past_key_values[layer_idx] if past_key_values is not None else None - raw = layer( - hidden_states, - rotary_pos_emb=rotary_pos_emb, - past_key_value=pkv, - attn_metadata=attn_metadata, - ) - raw_states.append(raw) - self.dd_encoder_self_dd.write_history(history_buffer, layer_idx + 1, raw) - hidden_states = self.dd_encoder_self_dd.forward_from_buffer(layer_idx, raw, history_buffer) - return hidden_states, raw_states - - def _decode(self, - decoder_input: torch.Tensor, - encoder_raw_states: list[torch.Tensor], - final_concept_state: torch.Tensor, - concept_raw_states: list[torch.Tensor] | None, - position_ids: torch.Tensor, - past_key_values: list[list[torch.Tensor]] | None = None, - attn_metadata: Any = None, - prefill_metadata: ConceptPrefillMetadata | None = None, - decode_concept_states: torch.Tensor | None = None): - """Decoder stack plus decoder DD and residual routes.""" - if decode_concept_states is not None: - decoder_encoder_states = torch.stack(tuple(encoder_raw_states), dim=-2) - decoder_encoder_states = self.decoder_read_encoder_shared_source_norm(decoder_encoder_states) - - concept_states = self.decoder_read_concept_shared_source_norm(decode_concept_states) - decoder_encoder_source_dim = -2 - repeated_concept_states = concept_states - elif prefill_metadata is None: - decoder_encoder_states = torch.stack(tuple(encoder_raw_states), dim=2) - decoder_encoder_states = self.decoder_read_encoder_shared_source_norm(decoder_encoder_states) - - concept_states = torch.stack(tuple(concept_raw_states), dim=2) - zero_chunk = torch.zeros_like(concept_states[:1]) - concept_states = torch.cat((zero_chunk, concept_states), dim=0) - concept_states = self.decoder_read_concept_shared_source_norm(concept_states) - decoder_encoder_source_dim = 2 - repeated_concept_states = None - else: - decoder_encoder_states = torch.stack(tuple(encoder_raw_states), dim=-2) - decoder_encoder_states = self.decoder_read_encoder_shared_source_norm(decoder_encoder_states) - - concept_states = torch.stack(tuple(concept_raw_states), dim=-2) - zero_chunk = torch.zeros_like(concept_states[:1]) - concept_states = torch.cat((zero_chunk, concept_states), dim=0) - concept_states = self.decoder_read_concept_shared_source_norm(concept_states) - decoder_encoder_source_dim = -2 - repeated_concept_states = self._repeat_shift_source_states_packed(concept_states, prefill_metadata) - - chunk_size = int(self.config.concept_chunk_size) - hidden_states = decoder_input - history_buffer = self.dd_two_route_add.make_history_buffer(hidden_states) - self.dd_two_route_add.write_history(history_buffer, 0, hidden_states) - rotary_pos_emb = self.decoder._make_rotary_pos_emb(hidden_states, position_ids) - - for layer_idx, layer in enumerate(self.decoder.layers): - pkv = past_key_values[layer_idx] if past_key_values is not None else None - raw = layer( - hidden_states, - rotary_pos_emb=rotary_pos_emb, - past_key_value=pkv, - attn_metadata=attn_metadata, - ) - self.dd_two_route_add.write_history(history_buffer, layer_idx + 1, raw) - gate = self._route_gate(layer_idx) - hidden_states = self.dd_two_route_add.forward_from_buffer( - layer_idx, - raw, - history_buffer, - final_concept_state, - gate[0], - ) - hidden_states = self.decoder_read_encoder_routes[layer_idx]( - hidden_states, - decoder_encoder_states, - source_dim=decoder_encoder_source_dim, - ) - if repeated_concept_states is None: - hidden_states = self.decoder_read_concept_routes[layer_idx].forward_repeated_chunks( - hidden_states, - concept_states, - chunk_size, - bool(getattr(self.config, 'concept_shift_feature', True)), - residual_scale=gate[1], - source_dim=2, - ) - else: - hidden_states = self.decoder_read_concept_routes[layer_idx]( - hidden_states, - repeated_concept_states, - residual_scale=gate[1], - source_dim=-2, - ) - - if self.decoder.final_layernorm is not None: - hidden_states = self.decoder.final_layernorm(hidden_states) - return hidden_states - - def prepare_inputs_for_generation(self, - past_key_values: list[list[torch.Tensor]], - inputs_embeds: torch.Tensor | None = None, - context: StepContext = None): - """Prepare input.""" - input_ids = context.input_ids - position_ids = context.position_ids - attn_metadata = context.attn_metadata - return dict( - input_ids=input_ids, - position_ids=position_ids, - past_key_values=past_key_values, - attn_metadata=attn_metadata, - inputs_embeds=inputs_embeds, - state_ids=context.state_offsets, - state_caches=context.state_caches, - ) - - def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]): - """Load native ConceptLM checkpoint weights into implemented - modules.""" - # (checkpoint_name, target_name) - weight_map = { - 'embedding.word_embeddings.weight': 'embedding.word_embeddings.weight', - 'output_layer.weight': 'lm_head.weight', - } - codebook_prefix = 'concept_quantizer.codebook.' - prediction_head_prefix = 'concept_predictor.prediction_heads.' - block_prefixes = ( - ('encoder.', self.encoder), - ('decoder.', self.decoder), - ('concept_predictor.hlm_block.', self.concept_predictor.hlm_block), - ) - params_dict = dict(self.named_parameters()) - for name, loaded_weight in weights: - if 'rotary_emb.inv_freq' in name: - continue - loaded_by_block = False - for block_prefix, block in block_prefixes: - if name.startswith(block_prefix): - block.load_weights([(name, loaded_weight)], prefix=block_prefix[:-1]) - loaded_by_block = True - break - if loaded_by_block: - continue - if name.startswith(prediction_head_prefix): - suffix = name[len(prediction_head_prefix):] - parts = suffix.split('.') - if len(parts) == 2 and parts[0].isdigit() and parts[1] in ('weight', 'bias'): - target_name = f'{prediction_head_prefix}proj.{parts[1]}' - param = params_dict.get(target_name) - if param is not None: - load_weight(param, loaded_weight, shard_id=int(parts[0])) - continue - if suffix in ('proj.weight', 'proj.bias'): - param = params_dict.get(name) - if param is not None: - for shard_id, shard_weight in enumerate(param.weight_spliter(loaded_weight)): - load_weight(param, shard_weight, shard_id=shard_id) - continue - if name.startswith(codebook_prefix): - codebook_idx = int(name[len(codebook_prefix):]) - param = params_dict['concept_quantizer.codebook'] - _load_stacked_codebook_weight(param, loaded_weight, codebook_idx) - continue - target = weight_map.get(name) - if target is None: - target = name - if target not in params_dict: - # skip checkpoint metadata or tensors owned by future runtime-only paths - continue - param = params_dict[target] - load_weight(param, loaded_weight) diff --git a/lmdeploy/pytorch/models/intern_ncp/__init__.py b/lmdeploy/pytorch/models/intern_ncp/__init__.py new file mode 100644 index 0000000000..b31d839e1e --- /dev/null +++ b/lmdeploy/pytorch/models/intern_ncp/__init__.py @@ -0,0 +1,6 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from .modeling import ConceptLMV22VQForCausalLM + +__all__ = [ + 'ConceptLMV22VQForCausalLM', +] diff --git a/lmdeploy/pytorch/models/intern_ncp/metadata.py b/lmdeploy/pytorch/models/intern_ncp/metadata.py new file mode 100644 index 0000000000..c7a55830a9 --- /dev/null +++ b/lmdeploy/pytorch/models/intern_ncp/metadata.py @@ -0,0 +1,278 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from collections.abc import Mapping +from dataclasses import dataclass +from typing import Any + +import torch +from transformers.configuration_utils import PretrainedConfig + +_CONCEPT_STATE_CHUNK_SOURCE_NAME = 'concept_chunk_source_state' +_CONCEPT_STATE_LAST_NAME = 'concept_last_state' +_CONCEPT_STATE_LAST_RAW_NAME = 'concept_last_raw_states' +_CONCEPT_STATE_LAST_FINAL_NAME = 'concept_last_final_state' + + +@dataclass +class ConceptMetadata: + """Layer-invariant ConceptLM runtime metadata. + + This mirrors the DSV4 pattern: ``StepContext`` is read at the top-level + model boundary, then submodules receive explicit metadata instead of + reaching back into the engine context. The dense/reference helpers below do + not consume all fields yet; they are part of the serving decode contract. + """ + + chunk_size: int + merge_method: str + shift_feature: bool + is_decoding: bool | None = None + state_ids: torch.Tensor | None = None + position_ids: torch.Tensor | None = None + block_offsets: torch.Tensor | None = None + q_seqlens: torch.Tensor | None = None + kv_seqlens: torch.Tensor | None = None + q_start_loc: torch.Tensor | None = None + attn_metadata: Any = None + + @classmethod + def build(cls, + config: PretrainedConfig, + position_ids: torch.Tensor | None = None, + attn_metadata: Any = None, + state_ids: torch.Tensor | None = None): + """Build ConceptLM metadata from explicit forward inputs.""" + return cls( + chunk_size=int(config.concept_chunk_size), + merge_method=getattr(config, 'concept_chunk_merge_method', 'meanpooling'), + shift_feature=bool(getattr(config, 'concept_shift_feature', True)), + is_decoding=getattr(attn_metadata, 'is_decoding', None), + state_ids=state_ids, + position_ids=position_ids, + block_offsets=getattr(attn_metadata, 'block_offsets', None), + q_seqlens=getattr(attn_metadata, 'q_seqlens', None), + kv_seqlens=getattr(attn_metadata, 'kv_seqlens', None), + q_start_loc=getattr(attn_metadata, 'q_start_loc', None), + attn_metadata=attn_metadata, + ) + + +@dataclass +class ConceptCaches: + """ConceptLM cache views resolved once at the top-level model boundary.""" + + encoder_past_key_values: list[list[torch.Tensor]] | None = None + concept_past_key_values: list[list[torch.Tensor]] | None = None + decoder_past_key_values: list[list[torch.Tensor]] | None = None + named_state_caches: Mapping[str, torch.Tensor] | None = None + state_caches: list[torch.Tensor] | None = None + chunk_source_name: str = _CONCEPT_STATE_CHUNK_SOURCE_NAME + last_state_name: str = _CONCEPT_STATE_LAST_NAME + last_raw_name: str = _CONCEPT_STATE_LAST_RAW_NAME + last_final_name: str = _CONCEPT_STATE_LAST_FINAL_NAME + chunk_source_idx: int = 0 + last_state_idx: int = 1 + last_raw_idx: int = 1 + last_final_idx: int = 2 + + @classmethod + def build(cls, + config: PretrainedConfig, + past_key_values: list[list[torch.Tensor]] | None = None, + state_caches: list[torch.Tensor] | None = None, + named_state_caches: Mapping[str, torch.Tensor] | None = None): + """Build ConceptLM cache views from engine-provided caches.""" + encoder_past_key_values, concept_past_key_values, decoder_past_key_values = ( + _split_concept_past_key_values(config, past_key_values)) + state_names = tuple(getattr(config, 'concept_state_names', ())) + + def _find_state_idx(state_name: str, fallback: int) -> int: + try: + return state_names.index(state_name) + except ValueError: + return fallback + + chunk_source_idx = int( + getattr(config, 'concept_state_chunk_source_idx', _find_state_idx(_CONCEPT_STATE_CHUNK_SOURCE_NAME, 0))) + last_state_idx = int(getattr(config, 'concept_state_last_idx', _find_state_idx(_CONCEPT_STATE_LAST_NAME, -1))) + last_raw_idx = int(getattr(config, 'concept_state_last_raw_idx', + _find_state_idx(_CONCEPT_STATE_LAST_RAW_NAME, 1))) + last_final_idx = int( + getattr(config, 'concept_state_last_final_idx', _find_state_idx(_CONCEPT_STATE_LAST_FINAL_NAME, 2))) + + def _state_name(state_idx: int, fallback: str) -> str: + if 0 <= state_idx < len(state_names): + return str(state_names[state_idx]) + return fallback + + return cls( + encoder_past_key_values=encoder_past_key_values, + concept_past_key_values=concept_past_key_values, + decoder_past_key_values=decoder_past_key_values, + named_state_caches=named_state_caches, + state_caches=state_caches, + chunk_source_name=_state_name(chunk_source_idx, _CONCEPT_STATE_CHUNK_SOURCE_NAME), + last_state_name=_state_name(last_state_idx, _CONCEPT_STATE_LAST_NAME), + last_raw_name=_state_name(last_raw_idx, _CONCEPT_STATE_LAST_RAW_NAME), + last_final_name=_state_name(last_final_idx, _CONCEPT_STATE_LAST_FINAL_NAME), + chunk_source_idx=chunk_source_idx, + last_state_idx=last_state_idx, + last_raw_idx=last_raw_idx, + last_final_idx=last_final_idx, + ) + + def named_state_cache(self, state_name: str) -> torch.Tensor | None: + """Return one named state-cache tensor when the engine provides it.""" + if self.named_state_caches is None or state_name not in self.named_state_caches: + return None + return self.named_state_caches[state_name] + + def state_cache(self, state_idx: int) -> torch.Tensor | None: + """Return one anonymous state-cache tensor by semantic index.""" + if self.state_caches is None: + return None + if state_idx < 0 or state_idx >= len(self.state_caches): + return None + return self.state_caches[state_idx] + + def semantic_state_cache(self, state_name: str, state_idx: int) -> torch.Tensor | None: + """Return a state cache by stable name, falling back to legacy + index.""" + cache = self.named_state_cache(state_name) + if cache is not None: + return cache + return self.state_cache(state_idx) + + @property + def chunk_source_state(self) -> torch.Tensor | None: + """Current chunk source accumulator state cache.""" + return self.semantic_state_cache(self.chunk_source_name, self.chunk_source_idx) + + @property + def last_state(self) -> torch.Tensor | None: + """Packed latest concept state cache. + + Shape is ``[num_state_slots, 1 + concept_layers, hidden]``. Row 0 is + the final concept vector; rows 1: are raw concept-layer states. + """ + return self.semantic_state_cache(self.last_state_name, self.last_state_idx) + + @property + def last_raw_states(self) -> torch.Tensor | None: + """Latest raw concept-layer state cache.""" + last_state = self.last_state + if last_state is not None: + return last_state[:, 1:] + return self.semantic_state_cache(self.last_raw_name, self.last_raw_idx) + + @property + def last_final_state(self) -> torch.Tensor | None: + """Latest final concept vector state cache.""" + last_state = self.last_state + if last_state is not None: + return last_state[:, 0] + return self.semantic_state_cache(self.last_final_name, self.last_final_idx) + + +@dataclass +class ConceptChunkStateUpdateResult: + """Fixed-shape result of one decode chunk-source state update.""" + + concept_input_states: torch.Tensor + concept_update_mask: torch.Tensor + + +@dataclass +class ConceptDecodeMetadata: + """Fixed-layout decode metadata derived once from engine inputs. + + Decode is always represented as the engine's fixed ``[1, batch]`` token + layout at the model boundary and flattened to ``[batch]`` / ``[batch, H]`` + only inside ConceptLM helpers. + """ + + position_ids: torch.Tensor + state_ids: torch.Tensor + safe_state_ids: torch.Tensor + valid_state_mask: torch.Tensor + + +@dataclass +class ConceptPrefillMetadata: + """Packed prefill metadata derived once from token attention metadata. + + Field groups: + - token stream: original engine-provided token request boundaries. + - concept stream: compact chunk-token request boundaries and positions + used by the concept predictor attention. + - chunk merge: token -> concept ids and helper ids used to reduce encoder + token states into concept states without a Python batch loop. + - repeat/gather: concept -> token ids used to project compact concept + states back to the packed token stream. + - scalar bounds: eager compact sizes / upper bounds needed by metadata and + attention launch parameters. + """ + + # Token stream metadata, shape [batch]. This is the original packed prefill + # layout consumed by normal token attention. + token_q_seqlens: torch.Tensor + token_q_start_loc: torch.Tensor + + # Concept stream metadata, shape [batch] plus compact concept positions. + # These describe the shorter chunk-token stream consumed by concept + # predictor attention. + concept_q_seqlens: torch.Tensor + concept_q_start_loc: torch.Tensor + concept_position_ids: torch.Tensor + + # Chunk merge metadata. ``merge_token_to_concept`` maps each packed token to + # the compact concept row that owns it, or -1 when the token is dropped from + # concept production. Counts/first/last ids implement mean/first/last merge. + merge_token_to_concept: torch.Tensor + merge_token_counts: torch.Tensor + merge_first_token_ids: torch.Tensor + merge_last_token_ids: torch.Tensor + merge_short_concept_mask: torch.Tensor + + # Repeat/gather metadata. Maps each packed token row to the compact concept + # row it should read after shift semantics are applied, or -1 for the + # zero-concept row. + token_to_concept: torch.Tensor + + # Scalar sizes/bounds. ``num_concepts_total`` is the exact compact size in + # eager prefill; ``max_concepts_per_request`` is the per-request attention + # launch bound. + num_tokens_total: int + num_concepts_total: int + max_concepts_per_request: int + + +def _split_concept_past_key_values(config: PretrainedConfig, past_key_values: list[list[torch.Tensor]] | None): + """Split the flat LMDeploy KV-cache list into ConceptLM streams.""" + if past_key_values is None or len(past_key_values) == 0: + return None, None, None + enc_layers = int(config.concept_encoder_layers) + concept_layers = int(config.concept_special_layers) + dec_layers = int(config.concept_decoder_layers) + total_layers = enc_layers + concept_layers + dec_layers + assert len(past_key_values) >= total_layers, ( + f'ConceptLM requires {total_layers} KV-cache layers ' + f'({enc_layers} encoder + {concept_layers} concept + {dec_layers} decoder), ' + f'got {len(past_key_values)}.') + enc_end = enc_layers + concept_end = enc_end + concept_layers + dec_end = concept_end + dec_layers + return past_key_values[:enc_end], past_key_values[enc_end:concept_end], past_key_values[concept_end:dec_end] + + +def _flatten_decode_position_ids(position_ids: torch.Tensor, batch_size: int) -> torch.Tensor: + """Normalize decode position ids to one absolute position per batch row.""" + if position_ids.dim() == 0: + position_ids = position_ids.view(1) + if position_ids.dim() == 1: + return position_ids.to(torch.long) + position_ids = position_ids.reshape(-1) + if position_ids.numel() == batch_size: + return position_ids.to(torch.long) + assert position_ids.numel() % batch_size == 0, ( + f'Cannot map position_ids with {position_ids.numel()} elements to batch size {batch_size}.') + return position_ids.reshape(-1, batch_size)[-1].to(torch.long) diff --git a/lmdeploy/pytorch/models/intern_ncp/modeling.py b/lmdeploy/pytorch/models/intern_ncp/modeling.py new file mode 100644 index 0000000000..ce6393c498 --- /dev/null +++ b/lmdeploy/pytorch/models/intern_ncp/modeling.py @@ -0,0 +1,1369 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from collections.abc import Iterable, Mapping +from dataclasses import dataclass, replace +from typing import Any + +import torch +from torch import nn +from torch.nn import functional as F +from transformers.configuration_utils import PretrainedConfig + +from lmdeploy.pytorch.model_inputs import StepContext, StepContextManager +from lmdeploy.pytorch.nn import ConceptLMRuntimeOps +from lmdeploy.pytorch.weight_loader.model_weight_loader import load_weight + +from ..patch import add_prefix +from ..utils.cudagraph import CudaGraphMixin +from ..utils.model import DeployModelMixinV1 +from .metadata import ( + ConceptCaches, + ConceptChunkStateUpdateResult, + ConceptDecodeMetadata, + ConceptMetadata, + ConceptPrefillMetadata, + _flatten_decode_position_ids, +) +from .modules import ( + ConceptPredictor, + Embedding, + OlmoBlock, + Quantizer, + ResidualRoute, + SelfDD, + TwoRouteAdd, +) +from .weight import _load_stacked_codebook_weight + + +@dataclass +class _ConceptPredictorOutput: + """Concept predictor output shared by prefill and decode paths.""" + + predicted_vectors: torch.Tensor + raw_states: list[torch.Tensor] + + +@dataclass +class _ConceptPredictorRequest: + """Concept-stream inputs consumed by the concept predictor.""" + + hidden_states: torch.Tensor + encoder_states: torch.Tensor + position_ids: torch.Tensor + attn_metadata: Any + + +@dataclass +class _DecoderConceptInput: + """Concept inputs consumed by the decoder stack.""" + + final_state: torch.Tensor + raw_states: list[torch.Tensor] | None = None + prefill_metadata: ConceptPrefillMetadata | None = None + decode_states: torch.Tensor | None = None + + @classmethod + def for_prefill(cls, + final_state: torch.Tensor, + raw_states: list[torch.Tensor], + prefill_metadata: ConceptPrefillMetadata): + """Build decoder concept inputs from compact prefill predictor + states.""" + return cls( + final_state=final_state, + raw_states=raw_states, + prefill_metadata=prefill_metadata, + ) + + @classmethod + def for_decode(cls, final_state: torch.Tensor, decode_states: torch.Tensor): + """Build decoder concept inputs from already gathered decode states.""" + return cls( + final_state=final_state, + decode_states=decode_states, + ) + + +@dataclass +class _EncoderOutput: + """Encoder result plus its reusable layer-major SelfDD history buffer.""" + + hidden_states: torch.Tensor + raw_states: list[torch.Tensor] + history_buffer: torch.Tensor + + +@dataclass +class _PrefillTokenLayout: + """Packed token-stream layout derived from prefill attention metadata.""" + + q_seqlens: torch.Tensor + q_start_loc: torch.Tensor + q_seqlens_long: torch.Tensor + q_start_loc_long: torch.Tensor + token_seq: torch.Tensor + token_pos: torch.Tensor + total_tokens: int + + +@dataclass +class _PrefillConceptLayout: + """Compact chunk-token stream layout used by concept predictor prefill.""" + + q_seqlens: torch.Tensor + q_seqlens_long: torch.Tensor + q_start_loc: torch.Tensor + q_start_loc_long: torch.Tensor + seq: torch.Tensor + local_ids: torch.Tensor + position_ids: torch.Tensor + num_total: int + max_per_request: int + + +@dataclass +class _PrefillMergeLayout: + """Token-to-concept merge metadata for compact prefill.""" + + token_to_concept: torch.Tensor + token_counts: torch.Tensor + first_token_ids: torch.Tensor + last_token_ids: torch.Tensor + short_concept_mask: torch.Tensor + + +class ConceptLMV22VQForCausalLM(nn.Module, DeployModelMixinV1, CudaGraphMixin): + """Rewrote model of ConceptLMV22VQForCausalLM.""" + + def __init__(self, + config: PretrainedConfig, + ctx_mgr: StepContextManager, + dtype: torch.dtype = None, + device: torch.device = None, + prefix: str = ''): + super().__init__() + self.config = config + self.ctx_mgr = ctx_mgr + # token embedding — mirrors ``self.embedding`` in the reference. + self.embedding = Embedding(config, dtype=dtype, device=device) + self.encoder = OlmoBlock(config, + config.concept_encoder_layers, + post_layer_norm=False, + dtype=dtype, + device=device, + prefix=add_prefix('encoder', prefix)) + self.decoder = OlmoBlock(config, + config.concept_decoder_layers, + post_layer_norm=True, + dtype=dtype, + device=device, + prefix=add_prefix('decoder', prefix)) + self.concept_vq_input_norm = nn.LayerNorm(config.hidden_size, + eps=getattr(config, 'layernorm_epsilon', 1e-6), + dtype=dtype, + device=device) + self.concept_quantizer = Quantizer(config, dtype=dtype, device=device) + self.concept_predictor = ConceptPredictor(config, + dtype=dtype, + device=device, + prefix=add_prefix('concept_predictor', prefix)) + self.fusion_tok_norm = nn.LayerNorm(config.hidden_size, + eps=getattr(config, 'layernorm_epsilon', 1e-6), + dtype=dtype, + device=device) + self.fusion_hl_norm = nn.LayerNorm(config.hidden_size, + eps=getattr(config, 'layernorm_epsilon', 1e-6), + dtype=dtype, + device=device) + self.fusion_norm_alpha = nn.Parameter( + torch.tensor(getattr(config, 'concept_fusion_norm_alpha_init', 0.1), dtype=dtype, device=device), + requires_grad=False) + self.dd_encoder_self_dd = SelfDD(config, + config.concept_encoder_layers, + use_softmax=False, + dtype=dtype, + device=device) + self.decoder_read_encoder_routes = nn.ModuleList([ + ResidualRoute(config, + config.concept_encoder_layers, + use_softmax=True, + dtype=dtype, + device=device) + for _ in range(config.concept_decoder_layers) + ]) + self.decoder_read_encoder_shared_source_norm = nn.LayerNorm(config.hidden_size, + eps=getattr(config, 'layernorm_epsilon', 1e-6), + dtype=dtype, + device=device) + self.decoder_read_concept_routes = nn.ModuleList([ + ResidualRoute(config, + config.concept_special_layers, + use_softmax=True, + dtype=dtype, + device=device) + for _ in range(config.concept_decoder_layers) + ]) + self.decoder_read_concept_shared_source_norm = nn.LayerNorm(config.hidden_size, + eps=getattr(config, 'layernorm_epsilon', 1e-6), + dtype=dtype, + device=device) + self.final_read_concept_gate_logits = nn.Parameter( + torch.zeros(config.concept_decoder_layers, 2, dtype=dtype, device=device), + requires_grad=False) + self.dd_two_route_add = TwoRouteAdd(config, dtype=dtype, device=device) + self.concept_ops = ConceptLMRuntimeOps() + # output projection — mirrors ``output_layer`` in the reference. Built + # via build_lm_head and named ``lm_head`` so DeployModelMixinV1. + # get_logits picks it up directly; load_weights maps the checkpoint's + # ``output_layer`` onto it. + self.lm_head = self.build_lm_head( + config.hidden_size, config.vocab_size, bias=False, dtype=dtype, device=device) + + def forward(self, + input_ids: torch.Tensor, + position_ids: torch.Tensor, + past_key_values: list[list[torch.Tensor]], + attn_metadata=None, + inputs_embeds: torch.Tensor = None, + state_ids: torch.Tensor | None = None, + state_caches: list[torch.Tensor] | None = None, + named_state_caches: Mapping[str, torch.Tensor] | None = None, + **kwargs): + """Model forward, return hidden_states (logits computed by runtime).""" + concept_metadata = self._build_concept_metadata(position_ids, attn_metadata, state_ids) + concept_caches = self._build_concept_caches(past_key_values, state_caches, named_state_caches) + if inputs_embeds is None: + hidden_states = self.embedding(input_ids) + else: + hidden_states = inputs_embeds + + if concept_metadata.is_decoding: + return self._forward_decode( + hidden_states, + position_ids, + concept_metadata, + concept_caches, + ) + + hidden_states, prefill_position_ids = self._normalize_prefill_inputs( + hidden_states, + position_ids, + attn_metadata, + ) + return self._forward_prefill_packed( + hidden_states, + prefill_position_ids, + concept_metadata, + concept_caches, + ) + + def get_input_embeddings(self): + """Get input embeddings.""" + return self.embedding.word_embeddings + + def get_output_embeddings(self): + """Get output embeddings.""" + return self.lm_head + + def _build_concept_metadata(self, + position_ids: torch.Tensor | None, + attn_metadata: Any = None, + state_ids: torch.Tensor | None = None) -> ConceptMetadata: + """Build top-level ConceptLM metadata for future serving paths.""" + return ConceptMetadata.build( + self.config, + position_ids=position_ids, + attn_metadata=attn_metadata, + state_ids=state_ids, + ) + + def _build_concept_caches(self, + past_key_values: list[list[torch.Tensor]] | None, + state_caches: list[torch.Tensor] | None = None, + named_state_caches: Mapping[str, torch.Tensor] | None = None) -> ConceptCaches: + """Build top-level ConceptLM cache views for future serving paths.""" + return ConceptCaches.build( + self.config, + past_key_values=past_key_values, + state_caches=state_caches, + named_state_caches=named_state_caches, + ) + + def _decode_chunk_state_update(self, + current_source_states: torch.Tensor, + concept_metadata: ConceptMetadata, + concept_caches: ConceptCaches) -> ConceptChunkStateUpdateResult: + """Update decode chunk-source state and return fixed-shape rows. + + CUDA uses the Triton writer. CPU uses the reference writer for tests. The returned rows deliberately avoid + dynamic concept-row compaction, matching the CUDA graph route in the design doc. + """ + chunk_source_state = concept_caches.chunk_source_state + concept_input_states, update_mask = self.concept_ops.decode_chunk_state_update( + chunk_source_state, + current_source_states, + concept_metadata.state_ids, + concept_metadata.position_ids, + concept_metadata.chunk_size, + concept_metadata.merge_method, + ) + return ConceptChunkStateUpdateResult( + concept_input_states=concept_input_states, + concept_update_mask=update_mask, + ) + + def _route_gate(self, layer_idx: int) -> torch.Tensor: + """Return decoder route gate ``[decoder_dd_scale, + concept_route_scale]``.""" + return self.final_read_concept_gate_logits[int(layer_idx)].float().softmax(dim=-1) + + def _normalize_prefill_inputs(self, + hidden_states: torch.Tensor, + position_ids: torch.Tensor, + attn_metadata: Any = None): + """Normalize LMDeploy prefill inputs to packed token layout. + + Engine layout stays fixed as ``[1, total_tokens, hidden]`` with + ``position_ids=[1, total_tokens]``. Per-request boundaries come from + ``attn_metadata.q_seqlens``/``q_start_loc`` and are used only to build + the concept stream. + """ + if attn_metadata is None or getattr(attn_metadata, 'q_seqlens', None) is None: + raise RuntimeError('ConceptLM prefill requires q_seqlens attention metadata.') + if getattr(attn_metadata, 'is_decoding', False): + raise RuntimeError('ConceptLM prefill input normalization received decode metadata.') + + if hidden_states.dim() != 3 or hidden_states.size(0) != 1: + raise NotImplementedError( + f'ConceptLM prefill expects fixed engine layout [1, total_tokens, hidden], ' + f'got {tuple(hidden_states.shape)}.') + + total_tokens = hidden_states.size(1) + if position_ids.dim() != 2 or position_ids.size(0) != 1: + raise NotImplementedError( + f'ConceptLM prefill expects fixed position layout [1, total_tokens], ' + f'got {tuple(position_ids.shape)}.') + if position_ids.size(1) != total_tokens: + raise ValueError(f'position_ids length {position_ids.size(1)} does not match token length {total_tokens}.') + + q_seqlens = attn_metadata.q_seqlens + if q_seqlens.dim() != 1: + raise ValueError(f'ConceptLM prefill expects 1-D q_seqlens, got {tuple(q_seqlens.shape)}.') + + hidden_states = hidden_states[0].contiguous() + return hidden_states, position_ids[0].to(device=hidden_states.device) + + def _normalize_decode_inputs(self, hidden_states: torch.Tensor, position_ids: torch.Tensor): + """Normalize LMDeploy decode inputs to one flat row per active + request.""" + if hidden_states.dim() != 3 or hidden_states.size(0) != 1: + raise NotImplementedError( + f'ConceptLM decode expects fixed engine layout [1, batch, hidden], ' + f'got {tuple(hidden_states.shape)}.') + batch_size = hidden_states.size(1) + position_ids = _flatten_decode_position_ids(position_ids, batch_size).to(device=hidden_states.device) + if position_ids.numel() != batch_size: + raise ValueError(f'Expected {batch_size} decode position ids, got {position_ids.numel()}.') + return hidden_states[0].contiguous(), position_ids + + @staticmethod + def _build_decode_metadata(position_ids: torch.Tensor, + state_ids: torch.Tensor | None, + batch_size: int, + device: torch.device) -> ConceptDecodeMetadata: + """Build fixed-shape decode metadata from engine state ids.""" + if state_ids is None: + raise RuntimeError('ConceptLM decode requires state_ids.') + state_ids = state_ids.to(device=device, dtype=torch.long).reshape(-1) + if state_ids.numel() != batch_size: + raise ValueError(f'Expected {batch_size} decode state ids, got {state_ids.numel()}.') + valid_state_mask = state_ids >= 0 + return ConceptDecodeMetadata( + position_ids=position_ids, + state_ids=state_ids, + safe_state_ids=state_ids.clamp(min=0), + valid_state_mask=valid_state_mask, + ) + + @staticmethod + def _select_decode_state_rows(state_cache: torch.Tensor, + decode_metadata: ConceptDecodeMetadata) -> torch.Tensor: + """Gather state-cache rows and zero out padded decode rows.""" + rows = state_cache.index_select(0, decode_metadata.safe_state_ids) + mask_shape = (decode_metadata.valid_state_mask.size(0), ) + (1, ) * (rows.dim() - 1) + valid_mask = decode_metadata.valid_state_mask.view(mask_shape) + return torch.where(valid_mask, rows, torch.zeros_like(rows)) + + def _select_decode_last_state_rows(self, concept_caches: ConceptCaches, + decode_metadata: ConceptDecodeMetadata) -> tuple[torch.Tensor, torch.Tensor]: + """Gather packed last-concept state rows once and return final/raw + views.""" + last_state = concept_caches.last_state + if last_state is not None: + rows = self._select_decode_state_rows(last_state, decode_metadata) + return rows[:, 0], rows[:, 1:] + + last_final_state = concept_caches.last_final_state + last_raw_states = concept_caches.last_raw_states + return ( + self._select_decode_state_rows(last_final_state, decode_metadata), + self._select_decode_state_rows(last_raw_states, decode_metadata), + ) + + def _decode_concept_read_mask(self, decode_metadata: ConceptDecodeMetadata) -> torch.Tensor: + """Return rows whose current decode token should read a cached + concept.""" + repeat_slots = self._repeat_slot_ids( + decode_metadata.position_ids, + int(self.config.concept_chunk_size), + bool(getattr(self.config, 'concept_shift_feature', True)), + ) + return decode_metadata.valid_state_mask & (repeat_slots >= 0) + + def _select_decode_decoder_concepts(self, + concept_caches: ConceptCaches, + decode_metadata: ConceptDecodeMetadata, + previous_final_state: torch.Tensor | None, + previous_raw_states: torch.Tensor | None) -> _DecoderConceptInput: + """Select the final/raw concept state visible to this decode token.""" + if bool(getattr(self.config, 'concept_shift_feature', True)): + final_state, raw_states = self._select_decode_last_state_rows(concept_caches, decode_metadata) + else: + final_state = previous_final_state + raw_states = previous_raw_states + + concept_read_mask = self._decode_concept_read_mask(decode_metadata) + final_state = torch.where(concept_read_mask.view(-1, 1), final_state, torch.zeros_like(final_state)) + raw_states = torch.where(concept_read_mask.view(-1, 1, 1), raw_states, torch.zeros_like(raw_states)) + return _DecoderConceptInput.for_decode(final_state, raw_states) + + @staticmethod + def _build_decode_chunk_source_states(encoder_output: _EncoderOutput) -> torch.Tensor: + """Return current per-row states accumulated until the next concept + boundary as a view over encoder history. + + Row 0 is the final encoder hidden, used as concept-predictor input when + a boundary is reached. Remaining rows mirror the prefill + ``encoder_raw_states[:-1]`` route sources. + """ + num_sources = max(len(encoder_output.raw_states), 1) + return encoder_output.history_buffer[:num_sources].movedim(0, -2) + + @staticmethod + def _decode_concept_position_ids(position_ids: torch.Tensor, chunk_size: int) -> torch.Tensor: + """Return reference RoPE positions for concept rows emitted at decode + boundaries.""" + return (position_ids - int(chunk_size) + 1).clamp(min=0) + + def _build_concept_decode_metadata_static(self, + token_attn_metadata: Any, + decode_metadata: ConceptDecodeMetadata): + """Build fixed-shape concept-stream decode metadata. + + Concept predictor runs with the same batch shape as token decode. + Boundary rows append a real concept KV entry. Non-boundary/padded rows + execute dummy concept attention with safe ``kv_seqlens >= 1``; their KV + writes are restored afterward by ``ConceptLMRuntimeOps``. + """ + device = decode_metadata.position_ids.device + batch_size = decode_metadata.position_ids.numel() + q_seqlens = token_attn_metadata.q_seqlens + q_start_loc = token_attn_metadata.q_start_loc + kv_seqlens = token_attn_metadata.kv_seqlens + q_dtype = q_seqlens.dtype + q_start_dtype = q_start_loc.dtype + kv_dtype = kv_seqlens.dtype + + concept_q_seqlens = torch.ones((batch_size, ), dtype=q_dtype, device=device) + concept_q_start_loc = torch.arange(batch_size, dtype=q_start_dtype, device=device) + concept_cu_seqlens = F.pad(torch.cumsum(concept_q_seqlens, dim=0, dtype=torch.int32), (1, 0)) + concept_kv_seqlens = torch.div( + decode_metadata.position_ids + 1, + int(self.config.concept_chunk_size), + rounding_mode='floor', + ).clamp(min=1).to(dtype=kv_dtype) + + updates = dict( + is_decoding=True, + block_offsets=token_attn_metadata.block_offsets, + q_start_loc=concept_q_start_loc, + q_seqlens=concept_q_seqlens, + kv_seqlens=concept_kv_seqlens, + cu_seqlens_q=concept_cu_seqlens, + cu_seqlens_k=concept_cu_seqlens, + ) + if hasattr(token_attn_metadata, 'kv_start_loc'): + updates['kv_start_loc'] = concept_kv_seqlens - concept_q_seqlens.to(dtype=concept_kv_seqlens.dtype) + if hasattr(token_attn_metadata, 'kv_flatten_size'): + updates['kv_flatten_size'] = batch_size + if hasattr(token_attn_metadata, 'max_q_seqlen'): + updates['max_q_seqlen'] = 1 + if hasattr(token_attn_metadata, 'max_kv_seqlen'): + # Decode kernels use kv_seqlens directly. Keep a conservative bound + # without reading the dynamic maximum back to host. + updates['max_kv_seqlen'] = getattr(token_attn_metadata, 'max_kv_seqlen') + if hasattr(token_attn_metadata, 'scheduler_metadata'): + updates['scheduler_metadata'] = None + if hasattr(token_attn_metadata, 'tile_scheduler_metadata'): + updates['tile_scheduler_metadata'] = None + if hasattr(token_attn_metadata, 'num_splits'): + updates['num_splits'] = None + if hasattr(token_attn_metadata, 'fill_seqlens'): + updates['fill_seqlens'] = concept_q_seqlens + return replace(token_attn_metadata, **updates) + + @staticmethod + def _stack_concept_raw_states(concept_raw_states: list[torch.Tensor]) -> torch.Tensor: + """Stack raw concept-layer states to ``[rows, concept_layers, + hidden]``.""" + return torch.stack(tuple(concept_raw_states), dim=1) + + def _snapshot_decode_concept_kv(self, + concept_caches: ConceptCaches, + concept_attn_metadata: Any): + """Snapshot concept KV slots that dummy non-boundary rows may + overwrite.""" + return [ + self.concept_ops.decode_kv_cache_snapshot( + k_cache, + v_cache, + concept_attn_metadata.block_offsets, + concept_attn_metadata.kv_seqlens, + ) + for k_cache, v_cache in concept_caches.concept_past_key_values + ] + + def _restore_decode_concept_kv_( + self, + concept_caches: ConceptCaches, + concept_attn_metadata: Any, + saved_kv, + restore_mask: torch.Tensor, + ): + """Restore concept KV slots for non-boundary and padded rows.""" + for (k_cache, v_cache), (saved_k, saved_v) in zip(concept_caches.concept_past_key_values, saved_kv): + self.concept_ops.decode_kv_cache_restore( + k_cache, + v_cache, + saved_k, + saved_v, + concept_attn_metadata.block_offsets, + concept_attn_metadata.kv_seqlens, + restore_mask, + ) + + def _write_decode_concept_states_static_(self, + concept_caches: ConceptCaches, + decode_metadata: ConceptDecodeMetadata, + update_mask: torch.Tensor, + predicted_vectors: torch.Tensor, + concept_raw_states: list[torch.Tensor]): + """Write newly emitted concept states to persistent decode caches.""" + last_final_state = concept_caches.last_final_state + last_raw_states = concept_caches.last_raw_states + raw_rows = self._stack_concept_raw_states(concept_raw_states) + self.concept_ops.decode_concept_state_update( + last_raw_states, + last_final_state, + predicted_vectors, + raw_rows, + decode_metadata.state_ids, + update_mask, + ) + + def _build_decode_concept_request(self, + chunk_update: ConceptChunkStateUpdateResult, + decode_metadata: ConceptDecodeMetadata, + concept_metadata: ConceptMetadata) -> _ConceptPredictorRequest: + """Build fixed-shape concept predictor inputs for decode.""" + concept_hidden = self.concept_vq_input_norm(chunk_update.concept_input_states[:, 0]) + encoder_concept_states = self.concept_predictor.normalize_encoder_concept_states( + chunk_update.concept_input_states[:, 1:]) + concept_position_ids = self._decode_concept_position_ids( + decode_metadata.position_ids, + concept_metadata.chunk_size, + ) + concept_attn_metadata = self._build_concept_decode_metadata_static( + concept_metadata.attn_metadata, + decode_metadata, + ) + return _ConceptPredictorRequest( + hidden_states=concept_hidden, + encoder_states=encoder_concept_states, + position_ids=concept_position_ids, + attn_metadata=concept_attn_metadata, + ) + + def _update_decode_concept_states_static_(self, + chunk_update: ConceptChunkStateUpdateResult, + decode_metadata: ConceptDecodeMetadata, + concept_metadata: ConceptMetadata, + concept_caches: ConceptCaches): + """Emit/cache concept states with fixed batch shape. + + The predictor runs for every decode row so CUDA graph capture sees a stable launch sequence. Non-boundary rows + are dummy work: their concept KV writes are restored and their final/raw state writes are masked. + """ + request = self._build_decode_concept_request(chunk_update, decode_metadata, concept_metadata) + saved_kv = self._snapshot_decode_concept_kv(concept_caches, request.attn_metadata) + concept_output = self._run_concept_predictor(request, concept_caches) + self._restore_decode_concept_kv_( + concept_caches, + request.attn_metadata, + saved_kv, + ~chunk_update.concept_update_mask, + ) + self._write_decode_concept_states_static_( + concept_caches, + decode_metadata, + chunk_update.concept_update_mask, + concept_output.predicted_vectors, + concept_output.raw_states, + ) + + def _merge_prefill_tail_chunk_states(self, + source_states: torch.Tensor, + prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: + """Build per-request partial chunk accumulator rows after prefill.""" + device = source_states.device + chunk_size = int(self.config.concept_chunk_size) + q_seqlens = prefill_metadata.token_q_seqlens.to(device=device, dtype=torch.long) + q_start_loc = prefill_metadata.token_q_start_loc.to(device=device, dtype=torch.long) + batch_size = q_seqlens.size(0) + tail_lens = torch.remainder(q_seqlens, chunk_size) + tail_lens = torch.where(q_seqlens < chunk_size, q_seqlens, tail_lens) + tail_lens = torch.where(q_seqlens > 0, tail_lens, torch.zeros_like(tail_lens)) + has_tail = tail_lens > 0 + + tail_rows = source_states.new_zeros((batch_size, source_states.size(1), source_states.size(2))) + if source_states.size(0) == 0: + return tail_rows + merge_method = getattr(self.config, 'concept_chunk_merge_method', 'meanpooling') + if merge_method == 'first': + token_ids = q_start_loc + q_seqlens - tail_lens + token_ids = token_ids.clamp(min=0, max=max(source_states.size(0) - 1, 0)) + rows = source_states[token_ids] + return torch.where(has_tail.view(batch_size, 1, 1), rows, tail_rows) + if merge_method == 'last': + token_ids = q_start_loc + q_seqlens - 1 + token_ids = token_ids.clamp(min=0, max=max(source_states.size(0) - 1, 0)) + rows = source_states[token_ids] + return torch.where(has_tail.view(batch_size, 1, 1), rows, tail_rows) + + token_ids = torch.arange(prefill_metadata.num_tokens_total, dtype=torch.long, device=device) + cu_q_seqlens = F.pad(torch.cumsum(q_seqlens, dim=0), (1, 0)) + token_seq = torch.searchsorted(cu_q_seqlens[1:], token_ids, right=True) + token_pos = token_ids - cu_q_seqlens[token_seq] + token_tail_start = q_seqlens[token_seq] - tail_lens[token_seq] + valid_tail = (tail_lens[token_seq] > 0) & (token_pos >= token_tail_start) + weighted_source = source_states * valid_tail.to(dtype=source_states.dtype).view(-1, 1, 1) + tail_rows.index_add_(0, token_seq, weighted_source) + return tail_rows + + def _write_prefill_state_caches_eager_(self, + concept_caches: ConceptCaches, + concept_metadata: ConceptMetadata, + prefill_metadata: ConceptPrefillMetadata, + source_states: torch.Tensor, + predicted_vectors: torch.Tensor, + concept_raw_states: list[torch.Tensor]): + """Seed decode state caches from a completed prefill forward.""" + if concept_metadata.state_ids is None: + return + if concept_caches.state_caches is None and concept_caches.named_state_caches is None: + return + chunk_source_state = concept_caches.chunk_source_state + last_raw_states = concept_caches.last_raw_states + last_final_state = concept_caches.last_final_state + + state_ids = concept_metadata.state_ids.to(device=source_states.device, dtype=torch.long).reshape(-1) + valid_indices = torch.nonzero(state_ids >= 0, as_tuple=False).flatten() + if valid_indices.numel() == 0: + return + + valid_state_ids = state_ids.index_select(0, valid_indices) + tail_rows = self._merge_prefill_tail_chunk_states(source_states, prefill_metadata) + chunk_source_state.index_copy_(0, valid_state_ids, + tail_rows.index_select(0, valid_indices).to(dtype=chunk_source_state.dtype)) + + concept_counts = prefill_metadata.concept_q_seqlens.to(device=source_states.device, dtype=torch.long) + concept_start = prefill_metadata.concept_q_start_loc.to(device=source_states.device, dtype=torch.long) + concept_valid_mask = (state_ids >= 0) & (concept_counts > 0) + concept_indices = torch.nonzero(concept_valid_mask, as_tuple=False).flatten() + if concept_indices.numel() == 0: + return + + concept_state_ids = state_ids.index_select(0, concept_indices) + last_concept_ids = concept_start.index_select(0, concept_indices) + concept_counts.index_select( + 0, concept_indices) - 1 + last_final_rows = predicted_vectors.index_select(0, last_concept_ids) + last_final_state.index_copy_(0, concept_state_ids, last_final_rows.to(dtype=last_final_state.dtype)) + raw_rows = self._stack_concept_raw_states(concept_raw_states).index_select(0, last_concept_ids) + last_raw_states.index_copy_(0, concept_state_ids, raw_rows.to(dtype=last_raw_states.dtype)) + + @staticmethod + def _concept_count_from_seq_len(seq_len: int, chunk_size: int) -> int: + """Return reference ConceptLM chunk count for one request length.""" + seq_len = int(seq_len or 0) + if seq_len <= 0: + return 0 + if seq_len < chunk_size: + return 1 + return seq_len // chunk_size + + @staticmethod + def _concept_counts_from_q_seqlens(q_seqlens: torch.Tensor, chunk_size: int) -> torch.Tensor: + """Vectorized version of ``_concept_count_from_seq_len``.""" + counts = torch.div(q_seqlens, chunk_size, rounding_mode='floor').clamp(min=1) + return torch.where(q_seqlens > 0, counts, torch.zeros_like(q_seqlens)) + + @staticmethod + def _repeat_slot_ids(token_pos: torch.Tensor, chunk_size: int, shift_feature: bool) -> torch.Tensor: + """Return local concept slot read by each token after shift + semantics.""" + if shift_feature: + return torch.div(token_pos + 1, chunk_size, rounding_mode='floor') - 1 + return torch.div(token_pos, chunk_size, rounding_mode='floor') - 1 + + @staticmethod + def _get_max_concepts_per_request(token_attn_metadata: Any, + concept_q_seqlens: torch.Tensor, + chunk_size: int) -> int: + """Return per-request concept attention bound without hidden context + access.""" + max_q_seqlen = getattr(token_attn_metadata, 'max_q_seqlen', None) + if max_q_seqlen is not None: + return ConceptLMV22VQForCausalLM._concept_count_from_seq_len(int(max_q_seqlen), chunk_size) + + # Test/direct-call fallback. Serving should use the scheduler-provided + # Python max_q_seqlen above, as in DSV4 metadata construction. + return int(concept_q_seqlens.max().item()) if concept_q_seqlens.numel() > 0 else 0 + + @staticmethod + def _build_prefill_token_layout(token_attn_metadata: Any, + position_ids: torch.Tensor) -> _PrefillTokenLayout: + """Build packed token-stream positions from prefill attention + metadata.""" + q_seqlens = token_attn_metadata.q_seqlens + q_start_loc = getattr(token_attn_metadata, 'q_start_loc', None) + if q_start_loc is None: + q_start_loc = F.pad(torch.cumsum(q_seqlens, dim=0, dtype=torch.int32), (1, 0))[:-1] + + total_tokens = int(position_ids.numel()) + + q_seqlens_long = q_seqlens.to(dtype=torch.long, device=position_ids.device) + q_start_loc_long = q_start_loc.to(dtype=torch.long, device=position_ids.device) + cu_q_seqlens = getattr(token_attn_metadata, 'cu_seqlens_q', None) + if cu_q_seqlens is None: + cu_q_seqlens = F.pad(torch.cumsum(q_seqlens, dim=0, dtype=torch.int32), (1, 0)) + cu_q_seqlens_long = cu_q_seqlens.to(dtype=torch.long, device=position_ids.device) + + token_ids = torch.arange(total_tokens, dtype=torch.long, device=position_ids.device) + token_seq = torch.searchsorted(cu_q_seqlens_long[1:], token_ids, right=True) + token_pos = token_ids - cu_q_seqlens_long[token_seq] + + return _PrefillTokenLayout( + q_seqlens=q_seqlens, + q_start_loc=q_start_loc, + q_seqlens_long=q_seqlens_long, + q_start_loc_long=q_start_loc_long, + token_seq=token_seq, + token_pos=token_pos, + total_tokens=total_tokens, + ) + + def _build_prefill_concept_layout(self, + token_attn_metadata: Any, + position_ids: torch.Tensor, + token_layout: _PrefillTokenLayout, + chunk_size: int) -> _PrefillConceptLayout: + """Build compact chunk-token positions used by concept attention.""" + concept_q_seqlens_long = self._concept_counts_from_q_seqlens(token_layout.q_seqlens_long, chunk_size) + concept_q_seqlens = concept_q_seqlens_long.to(dtype=token_layout.q_seqlens.dtype, + device=token_layout.q_seqlens.device) + concept_q_start_loc = F.pad(torch.cumsum(concept_q_seqlens, dim=0, dtype=torch.int32), (1, 0))[:-1] + concept_q_start_loc_long = concept_q_start_loc.to(dtype=torch.long, device=position_ids.device) + concept_cu_seqlens_long = F.pad(torch.cumsum(concept_q_seqlens_long, dim=0), (1, 0)) + + # TODO: remove this eager compact-size scalar read by moving ConceptLM + # concept-stream allocation/compaction into the engine/backend contract, + # like DSV4's precomputed metadata and Qwen3.5's state-cache metadata. + num_concepts_total = int(concept_cu_seqlens_long[-1].item()) + concept_ids = torch.arange(num_concepts_total, dtype=torch.long, device=position_ids.device) + concept_seq = torch.searchsorted(concept_cu_seqlens_long[1:], concept_ids, right=True) + local_concept_ids = concept_ids - concept_cu_seqlens_long[concept_seq] + concept_token_start = token_layout.q_start_loc_long[concept_seq] + local_concept_ids * chunk_size + concept_position_ids = position_ids[concept_token_start] + max_concepts_per_request = self._get_max_concepts_per_request( + token_attn_metadata, + concept_q_seqlens_long, + chunk_size, + ) + + return _PrefillConceptLayout( + q_seqlens=concept_q_seqlens, + q_seqlens_long=concept_q_seqlens_long, + q_start_loc=concept_q_start_loc, + q_start_loc_long=concept_q_start_loc_long, + seq=concept_seq, + local_ids=local_concept_ids, + position_ids=concept_position_ids, + num_total=num_concepts_total, + max_per_request=max_concepts_per_request, + ) + + def _build_prefill_repeat_ids(self, + token_layout: _PrefillTokenLayout, + concept_layout: _PrefillConceptLayout, + chunk_size: int, + shift_feature: bool) -> torch.Tensor: + """Map packed token rows to the concept row visible after shift.""" + seq_concept_start = concept_layout.q_start_loc_long[token_layout.token_seq] + seq_concept_count = concept_layout.q_seqlens_long[token_layout.token_seq] + repeat_slots = self._repeat_slot_ids(token_layout.token_pos, chunk_size, shift_feature) + valid_repeat = (repeat_slots >= 0) & (repeat_slots < seq_concept_count) + return torch.where( + valid_repeat, + seq_concept_start + repeat_slots, + torch.full_like(repeat_slots, -1), + ) + + @staticmethod + def _build_prefill_merge_layout(token_layout: _PrefillTokenLayout, + concept_layout: _PrefillConceptLayout, + chunk_size: int) -> _PrefillMergeLayout: + """Build token-to-concept merge metadata for mean/first/last modes.""" + seq_concept_start = concept_layout.q_start_loc_long[token_layout.token_seq] + seq_concept_count = concept_layout.q_seqlens_long[token_layout.token_seq] + merge_slots = torch.div(token_layout.token_pos, chunk_size, rounding_mode='floor') + valid_merge = (merge_slots >= 0) & (merge_slots < seq_concept_count) + merge_token_to_concept = torch.where( + valid_merge, + seq_concept_start + merge_slots, + torch.full_like(merge_slots, -1), + ) + safe_merge_ids = merge_token_to_concept.clamp(min=0) + merge_token_counts = torch.zeros(concept_layout.num_total, + dtype=torch.int32, + device=token_layout.token_pos.device) + merge_token_counts.index_add_(0, safe_merge_ids, valid_merge.to(dtype=torch.int32)) + + concept_seq_len = token_layout.q_seqlens_long[concept_layout.seq] + merge_first_pos = torch.where( + concept_seq_len < chunk_size, + torch.zeros_like(concept_layout.local_ids), + concept_layout.local_ids * chunk_size, + ) + merge_last_pos = torch.where( + concept_seq_len < chunk_size, + (concept_seq_len - 1).clamp(min=0), + concept_layout.local_ids * chunk_size + chunk_size - 1, + ) + merge_first_token_ids = token_layout.q_start_loc_long[concept_layout.seq] + merge_first_pos + merge_last_token_ids = token_layout.q_start_loc_long[concept_layout.seq] + merge_last_pos + merge_short_concept_mask = concept_seq_len < chunk_size + + return _PrefillMergeLayout( + token_to_concept=merge_token_to_concept, + token_counts=merge_token_counts, + first_token_ids=merge_first_token_ids, + last_token_ids=merge_last_token_ids, + short_concept_mask=merge_short_concept_mask, + ) + + def _build_prefill_metadata(self, + token_attn_metadata: Any, + position_ids: torch.Tensor) -> ConceptPrefillMetadata: + """Build packed token-to-concept metadata for batched prefill.""" + chunk_size = int(self.config.concept_chunk_size) + shift_feature = bool(getattr(self.config, 'concept_shift_feature', True)) + token_layout = self._build_prefill_token_layout(token_attn_metadata, position_ids) + concept_layout = self._build_prefill_concept_layout( + token_attn_metadata, + position_ids, + token_layout, + chunk_size, + ) + token_to_concept = self._build_prefill_repeat_ids( + token_layout, + concept_layout, + chunk_size, + shift_feature, + ) + merge_layout = self._build_prefill_merge_layout(token_layout, concept_layout, chunk_size) + + return ConceptPrefillMetadata( + token_q_seqlens=token_layout.q_seqlens, + token_q_start_loc=token_layout.q_start_loc, + concept_q_seqlens=concept_layout.q_seqlens, + concept_q_start_loc=concept_layout.q_start_loc, + concept_position_ids=concept_layout.position_ids, + merge_token_to_concept=merge_layout.token_to_concept, + merge_token_counts=merge_layout.token_counts, + merge_first_token_ids=merge_layout.first_token_ids, + merge_last_token_ids=merge_layout.last_token_ids, + merge_short_concept_mask=merge_layout.short_concept_mask, + token_to_concept=token_to_concept, + num_tokens_total=token_layout.total_tokens, + num_concepts_total=concept_layout.num_total, + max_concepts_per_request=concept_layout.max_per_request, + ) + + @staticmethod + def _merge_chunks_mean_packed(hidden_states: torch.Tensor, + prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: + """Mean-pool packed token states by precomputed concept ids.""" + merge_token_to_concept = prefill_metadata.merge_token_to_concept.to(device=hidden_states.device) + valid_merge = (merge_token_to_concept >= 0).to(dtype=hidden_states.dtype).unsqueeze(-1) + safe_merge_ids = merge_token_to_concept.clamp(min=0) + merged = hidden_states.new_zeros((prefill_metadata.num_concepts_total, hidden_states.size(-1))) + merged.index_add_(0, safe_merge_ids, hidden_states * valid_merge) + counts = prefill_metadata.merge_token_counts.clamp(min=1).to(device=hidden_states.device, + dtype=hidden_states.dtype) + return merged / counts.unsqueeze(-1) + + def _merge_chunks_packed(self, + hidden_states: torch.Tensor, + prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: + """Merge packed token states into packed per-request concept states.""" + merge_method = getattr(self.config, 'concept_chunk_merge_method', 'meanpooling') + if merge_method == 'first': + merged = hidden_states[prefill_metadata.merge_first_token_ids] + short_mean = self._merge_chunks_mean_packed(hidden_states, prefill_metadata) + return torch.where(prefill_metadata.merge_short_concept_mask[:, None], short_mean, merged) + if merge_method == 'last': + merged = hidden_states[prefill_metadata.merge_last_token_ids] + short_mean = self._merge_chunks_mean_packed(hidden_states, prefill_metadata) + return torch.where(prefill_metadata.merge_short_concept_mask[:, None], short_mean, merged) + + return self._merge_chunks_mean_packed(hidden_states, prefill_metadata) + + @staticmethod + def _gather_zero_prefixed_concepts(concept_states_with_zero: torch.Tensor, + prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: + """Gather zero-prefixed concept rows to packed token rows.""" + token_to_concept = prefill_metadata.token_to_concept.to(device=concept_states_with_zero.device) + gather_ids = torch.clamp(token_to_concept + 1, min=0) + return concept_states_with_zero[gather_ids] + + def _repeat_shift_packed(self, + concept_states: torch.Tensor, + prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: + """Gather packed concept states back to packed token states.""" + concept_states_with_zero = torch.cat((torch.zeros_like(concept_states[:1]), concept_states), dim=0) + return self._gather_zero_prefixed_concepts(concept_states_with_zero, prefill_metadata) + + def _repeat_shift_source_states_packed(self, + concept_states_with_zero: torch.Tensor, + prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: + """Gather zero-prefixed packed concept source states to token rows.""" + return self._gather_zero_prefixed_concepts(concept_states_with_zero, prefill_metadata) + + def _build_encoder_concept_states_packed(self, + encoder_raw_states: list[torch.Tensor], + prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: + """Build packed chunk-level encoder states used by the concept + predictor.""" + chunks = [self._merge_chunks_packed(state, prefill_metadata) for state in encoder_raw_states[:-1]] + states = torch.stack(chunks, dim=-2) + return self.concept_predictor.normalize_encoder_concept_states(states) + + def _run_concept_predictor(self, + request: _ConceptPredictorRequest, + concept_caches: ConceptCaches) -> _ConceptPredictorOutput: + """Run concept predictor and quantizer for any concept stream + layout.""" + concept_logits, concept_raw_states = self.concept_predictor( + request.hidden_states, + request.encoder_states, + request.position_ids, + past_key_values=concept_caches.concept_past_key_values, + attn_metadata=request.attn_metadata, + ) + return _ConceptPredictorOutput( + predicted_vectors=self.concept_quantizer(concept_logits), + raw_states=concept_raw_states, + ) + + def _build_prefill_concept_request(self, + hidden_states: torch.Tensor, + encoder_raw_states: list[torch.Tensor], + prefill_metadata: ConceptPrefillMetadata, + concept_metadata: ConceptMetadata) -> _ConceptPredictorRequest: + """Build concept predictor inputs for compact packed prefill.""" + concept_hidden = self._merge_chunks_packed(hidden_states, prefill_metadata) + concept_hidden = self.concept_vq_input_norm(concept_hidden) + encoder_concept_states = self._build_encoder_concept_states_packed(encoder_raw_states, prefill_metadata) + concept_attn_metadata = self._build_concept_prefill_metadata( + concept_metadata.attn_metadata, + prefill_metadata, + ) + return _ConceptPredictorRequest( + hidden_states=concept_hidden, + encoder_states=encoder_concept_states, + position_ids=prefill_metadata.concept_position_ids, + attn_metadata=concept_attn_metadata, + ) + + def _fuse_token_concept_states(self, + hidden_states: torch.Tensor, + final_concept_state: torch.Tensor) -> torch.Tensor: + """Fuse token and final-concept states before the decoder stack.""" + return self.fusion_tok_norm(hidden_states) + self.fusion_norm_alpha.to( + hidden_states.dtype) * self.fusion_hl_norm(final_concept_state.to(hidden_states.dtype)) + + def _run_decoder_from_concepts(self, + hidden_states: torch.Tensor, + encoder_raw_states: list[torch.Tensor], + decoder_concepts: _DecoderConceptInput, + position_ids: torch.Tensor, + concept_caches: ConceptCaches, + attn_metadata: Any) -> torch.Tensor: + """Fuse visible concept state and run the decoder stack.""" + decoder_input = self._fuse_token_concept_states(hidden_states, decoder_concepts.final_state) + return self._decode( + decoder_input, + encoder_raw_states, + decoder_concepts, + position_ids, + concept_caches, + attn_metadata, + ) + + def _prepare_decoder_route_sources(self, + encoder_raw_states: list[torch.Tensor], + decoder_concepts: _DecoderConceptInput) -> tuple[torch.Tensor, torch.Tensor]: + """Prepare encoder/concept route sources in packed ``[..., L, H]`` + layout.""" + decoder_encoder_states = torch.stack(tuple(encoder_raw_states), dim=-2) + decoder_encoder_states = self.decoder_read_encoder_shared_source_norm(decoder_encoder_states) + + if decoder_concepts.decode_states is not None: + concept_states = self.decoder_read_concept_shared_source_norm(decoder_concepts.decode_states) + return decoder_encoder_states, concept_states + + concept_states = torch.stack(tuple(decoder_concepts.raw_states), dim=-2) + zero_chunk = torch.zeros_like(concept_states[:1]) + concept_states = torch.cat((zero_chunk, concept_states), dim=0) + concept_states = self.decoder_read_concept_shared_source_norm(concept_states) + concept_states = self._repeat_shift_source_states_packed(concept_states, decoder_concepts.prefill_metadata) + return decoder_encoder_states, concept_states + + def _build_concept_prefill_metadata(self, + token_attn_metadata: Any, + prefill_metadata: ConceptPrefillMetadata): + """Build chunk-stream attention metadata for packed prefill.""" + concept_q_seqlens = prefill_metadata.concept_q_seqlens + concept_q_start_loc = prefill_metadata.concept_q_start_loc + concept_cu_seqlens = torch.nn.functional.pad( + torch.cumsum(concept_q_seqlens, dim=0, dtype=torch.int32), + (1, 0), + ) + max_concept_seqlen = int(prefill_metadata.max_concepts_per_request) + + updates = dict( + is_decoding=False, + q_start_loc=concept_q_start_loc, + q_seqlens=concept_q_seqlens, + kv_seqlens=concept_q_seqlens, + cu_seqlens_q=concept_cu_seqlens, + cu_seqlens_k=concept_cu_seqlens, + ) + if hasattr(token_attn_metadata, 'kv_start_loc'): + updates['kv_start_loc'] = concept_q_start_loc + if hasattr(token_attn_metadata, 'kv_flatten_size'): + updates['kv_flatten_size'] = int(prefill_metadata.num_concepts_total) + if hasattr(token_attn_metadata, 'max_q_seqlen'): + updates['max_q_seqlen'] = max_concept_seqlen + if hasattr(token_attn_metadata, 'max_kv_seqlen'): + updates['max_kv_seqlen'] = max_concept_seqlen + return replace(token_attn_metadata, **updates) + + def _forward_prefill_packed(self, + hidden_states: torch.Tensor, + position_ids: torch.Tensor, + concept_metadata: ConceptMetadata, + concept_caches: ConceptCaches): + """Packed non-decode ConceptLM forward.""" + if (concept_caches.encoder_past_key_values is None or concept_caches.concept_past_key_values is None + or concept_caches.decoder_past_key_values is None): + raise RuntimeError('ConceptLM prefill requires encoder, concept, and decoder KV caches.') + prefill_metadata = self._build_prefill_metadata(concept_metadata.attn_metadata, position_ids) + encoder_output = self._encode( + hidden_states, + position_ids, + past_key_values=concept_caches.encoder_past_key_values, + attn_metadata=concept_metadata.attn_metadata, + ) + hidden_states = encoder_output.hidden_states + encoder_raw_states = encoder_output.raw_states + concept_request = self._build_prefill_concept_request( + hidden_states, + encoder_raw_states, + prefill_metadata, + concept_metadata, + ) + concept_output = self._run_concept_predictor(concept_request, concept_caches) + self._write_prefill_state_caches_eager_( + concept_caches, + concept_metadata, + prefill_metadata, + self._build_decode_chunk_source_states(encoder_output), + concept_output.predicted_vectors, + concept_output.raw_states, + ) + repeated_concepts = self._repeat_shift_packed(concept_output.predicted_vectors, prefill_metadata) + decoder_concepts = _DecoderConceptInput.for_prefill( + repeated_concepts, + concept_output.raw_states, + prefill_metadata, + ) + final_hidden = self._run_decoder_from_concepts( + hidden_states, + encoder_raw_states, + decoder_concepts, + position_ids, + concept_caches, + concept_metadata.attn_metadata, + ) + return final_hidden.unsqueeze(0).contiguous() + + def _forward_decode(self, + hidden_states: torch.Tensor, + position_ids: torch.Tensor, + concept_metadata: ConceptMetadata, + concept_caches: ConceptCaches): + """ConceptLM decode path. + + This path is semantically structured for serving. Boundary concept + updates run with fixed batch shape so it is eligible for CUDA graph + replay through the base ``CudaGraphMixin`` decode-only policy. + """ + if (concept_caches.encoder_past_key_values is None or concept_caches.concept_past_key_values is None + or concept_caches.decoder_past_key_values is None): + raise RuntimeError('ConceptLM decode requires encoder, concept, and decoder KV caches.') + if concept_caches.chunk_source_state is None: + raise RuntimeError('ConceptLM decode requires chunk source state cache.') + if concept_caches.last_final_state is None or concept_caches.last_raw_states is None: + raise RuntimeError('ConceptLM decode requires cached last concept states.') + + hidden_states, decode_position_ids = self._normalize_decode_inputs(hidden_states, position_ids) + decode_metadata = self._build_decode_metadata( + decode_position_ids, + concept_metadata.state_ids, + hidden_states.size(0), + hidden_states.device, + ) + decode_concept_metadata = replace( + concept_metadata, + position_ids=decode_position_ids, + state_ids=decode_metadata.state_ids, + ) + + encoder_output = self._encode( + hidden_states, + decode_position_ids, + past_key_values=concept_caches.encoder_past_key_values, + attn_metadata=concept_metadata.attn_metadata, + ) + hidden_states = encoder_output.hidden_states + encoder_raw_states = encoder_output.raw_states + previous_final_concept_state = None + previous_concept_raw_state_rows = None + if not bool(getattr(self.config, 'concept_shift_feature', True)): + previous_final_concept_state, previous_concept_raw_state_rows = self._select_decode_last_state_rows( + concept_caches, decode_metadata) + current_source_states = self._build_decode_chunk_source_states(encoder_output) + chunk_update = self._decode_chunk_state_update( + current_source_states, + decode_concept_metadata, + concept_caches, + ) + self._update_decode_concept_states_static_( + chunk_update, + decode_metadata, + decode_concept_metadata, + concept_caches, + ) + + decoder_concepts = self._select_decode_decoder_concepts( + concept_caches, + decode_metadata, + previous_final_concept_state, + previous_concept_raw_state_rows, + ) + final_hidden = self._run_decoder_from_concepts( + hidden_states, + encoder_raw_states, + decoder_concepts, + decode_position_ids, + concept_caches, + concept_metadata.attn_metadata, + ) + return final_hidden.unsqueeze(0).contiguous() + + def _encode(self, + hidden_states: torch.Tensor, + position_ids: torch.Tensor, + past_key_values: list[list[torch.Tensor]] | None = None, + attn_metadata: Any = None): + """Encoder stack plus encoder self-DD. + + This helper mirrors the reference flow but consumes LMDeploy attention inputs. It is wired for the future full + forward path; continuous batching still needs caller-side concept metadata before the full model can use it + safely. + """ + raw_states = [] + history_buffer = self.dd_encoder_self_dd.make_history_buffer(hidden_states) + self.dd_encoder_self_dd.write_history(history_buffer, 0, hidden_states) + rotary_pos_emb = self.encoder._make_rotary_pos_emb(hidden_states, position_ids) + + for layer_idx, layer in enumerate(self.encoder.layers): + pkv = past_key_values[layer_idx] if past_key_values is not None else None + raw = layer( + hidden_states, + rotary_pos_emb=rotary_pos_emb, + past_key_value=pkv, + attn_metadata=attn_metadata, + ) + raw_states.append(raw) + self.dd_encoder_self_dd.write_history(history_buffer, layer_idx + 1, raw) + hidden_states = self.dd_encoder_self_dd.forward_from_buffer(layer_idx, raw, history_buffer) + + # The encoder SelfDD path no longer needs the original input stored in + # slot 0 after the loop. Reuse that slot for the final encoder hidden + # so decode/prefill state seeding can view ``[final, raw[:-1]]`` + # without stacking and copying every raw layer output again. + self.dd_encoder_self_dd.write_history(history_buffer, 0, hidden_states) + return _EncoderOutput( + hidden_states=hidden_states, + raw_states=raw_states, + history_buffer=history_buffer, + ) + + def _decode(self, + decoder_input: torch.Tensor, + encoder_raw_states: list[torch.Tensor], + decoder_concepts: _DecoderConceptInput, + position_ids: torch.Tensor, + concept_caches: ConceptCaches, + attn_metadata: Any = None): + """Decoder stack plus decoder DD and residual routes.""" + past_key_values = concept_caches.decoder_past_key_values + decoder_encoder_states, concept_states = self._prepare_decoder_route_sources( + encoder_raw_states, + decoder_concepts, + ) + final_concept_state = decoder_concepts.final_state + hidden_states = decoder_input + history_buffer = self.dd_two_route_add.make_history_buffer(hidden_states) + self.dd_two_route_add.write_history(history_buffer, 0, hidden_states) + rotary_pos_emb = self.decoder._make_rotary_pos_emb(hidden_states, position_ids) + + for layer_idx, layer in enumerate(self.decoder.layers): + pkv = past_key_values[layer_idx] if past_key_values is not None else None + raw = layer( + hidden_states, + rotary_pos_emb=rotary_pos_emb, + past_key_value=pkv, + attn_metadata=attn_metadata, + ) + self.dd_two_route_add.write_history(history_buffer, layer_idx + 1, raw) + gate = self._route_gate(layer_idx) + hidden_states = self.dd_two_route_add.forward_from_buffer( + layer_idx, + raw, + history_buffer, + final_concept_state, + gate[0], + ) + hidden_states = self.decoder_read_encoder_routes[layer_idx]( + hidden_states, + decoder_encoder_states, + source_dim=-2, + ) + hidden_states = self.decoder_read_concept_routes[layer_idx]( + hidden_states, + concept_states, + residual_scale=gate[1], + source_dim=-2, + ) + + if self.decoder.final_layernorm is not None: + hidden_states = self.decoder.final_layernorm(hidden_states) + return hidden_states + + def prepare_inputs_for_generation(self, + past_key_values: list[list[torch.Tensor]], + inputs_embeds: torch.Tensor | None = None, + context: StepContext = None): + """Prepare input.""" + input_ids = context.input_ids + position_ids = context.position_ids + attn_metadata = context.attn_metadata + return dict( + input_ids=input_ids, + position_ids=position_ids, + past_key_values=past_key_values, + attn_metadata=attn_metadata, + inputs_embeds=inputs_embeds, + state_ids=context.state_offsets, + state_caches=context.state_caches, + named_state_caches=context.named_state_caches, + ) + + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]): + """Load native ConceptLM checkpoint weights into implemented + modules.""" + # (checkpoint_name, target_name) + weight_map = { + 'embedding.word_embeddings.weight': 'embedding.word_embeddings.weight', + 'output_layer.weight': 'lm_head.weight', + } + codebook_prefix = 'concept_quantizer.codebook.' + prediction_head_prefix = 'concept_predictor.prediction_heads.' + block_prefixes = ( + ('encoder.', self.encoder), + ('decoder.', self.decoder), + ('concept_predictor.hlm_block.', self.concept_predictor.hlm_block), + ) + params_dict = dict(self.named_parameters()) + for name, loaded_weight in weights: + if 'rotary_emb.inv_freq' in name: + continue + loaded_by_block = False + for block_prefix, block in block_prefixes: + if name.startswith(block_prefix): + block.load_weights([(name, loaded_weight)], prefix=block_prefix[:-1]) + loaded_by_block = True + break + if loaded_by_block: + continue + if name.startswith(prediction_head_prefix): + suffix = name[len(prediction_head_prefix):] + parts = suffix.split('.') + if len(parts) == 2 and parts[0].isdigit() and parts[1] in ('weight', 'bias'): + target_name = f'{prediction_head_prefix}proj.{parts[1]}' + param = params_dict.get(target_name) + if param is not None: + load_weight(param, loaded_weight, shard_id=int(parts[0])) + continue + if suffix in ('proj.weight', 'proj.bias'): + param = params_dict.get(name) + if param is not None: + for shard_id, shard_weight in enumerate(param.weight_spliter(loaded_weight)): + load_weight(param, shard_weight, shard_id=shard_id) + continue + if name.startswith(codebook_prefix): + codebook_idx = int(name[len(codebook_prefix):]) + param = params_dict['concept_quantizer.codebook'] + _load_stacked_codebook_weight(param, loaded_weight, codebook_idx) + continue + target = weight_map.get(name) + if target is None: + target = name + if target not in params_dict: + # skip checkpoint metadata or tensors owned by future runtime-only paths + continue + param = params_dict[target] + load_weight(param, loaded_weight) diff --git a/lmdeploy/pytorch/models/intern_ncp/modules.py b/lmdeploy/pytorch/models/intern_ncp/modules.py new file mode 100644 index 0000000000..7e77eae76f --- /dev/null +++ b/lmdeploy/pytorch/models/intern_ncp/modules.py @@ -0,0 +1,984 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from collections.abc import Iterable +from typing import Any + +import torch +from torch import nn +from torch.nn import functional as F +from transformers.configuration_utils import PretrainedConfig + +from lmdeploy.pytorch.nn import ( + ApplyRotaryEmb, + Attention, + RMSNorm, + SiluAndMul, + build_rotary_embedding, +) +from lmdeploy.pytorch.nn.linear import ( + build_down_linear, + build_gateup_linear, + build_merged_colwise_linear, + build_o_proj, + build_qkv_proj, +) +from lmdeploy.pytorch.weight_loader.model_weight_loader import load_weight + +from ..patch import add_prefix +from ..utils.model import build_embedding +from .weight import _repack_olmo_qkv_weight + +_CONFIG_VALUE = object() +_HistoryStates = torch.Tensor +_SourceStates = torch.Tensor | None + +def _get_configured_window(config: PretrainedConfig): + """Return the reference OLMo window setting as an int or None.""" + window_size = getattr(config, 'window_size', None) + if window_size is None: + return None + if isinstance(window_size, (list, tuple)): + window_size = window_size[0] + if window_size is None: + return None + window_size = int(window_size) + return window_size if window_size > 0 else None + + +def _make_olmo_rotary_embedding(config: PretrainedConfig, + device: torch.device = None) -> nn.Module: + """Build ConceptLM/OLMo rotary embedding from its native config fields.""" + rotary_interleaved = bool(getattr(config, 'rotary_interleaved', False)) + if rotary_interleaved: + raise NotImplementedError('ConceptLM rotary_interleaved=True is not supported by the LMDeploy block yet.') + + head_dim = int(config.kv_channels) + rotary_percent = float(getattr(config, 'rotary_percent', 1.0)) + rotary_dim = int(head_dim * rotary_percent) + rotary_dim -= rotary_dim % 2 + if rotary_dim <= 0: + raise ValueError(f'Invalid ConceptLM rotary dimension: head_dim={head_dim}, rotary_percent={rotary_percent}') + + partial_rotary_factor = rotary_dim / head_dim + return build_rotary_embedding( + dim=head_dim, + max_position_embeddings=getattr(config, 'max_position_embeddings', getattr(config, 'max_sequence_length', + 2048)), + base=getattr(config, 'rotary_base', 10000), + partial_rotary_factor=partial_rotary_factor, + device=device, + ) + +def _qk_rmsnorm_variance(query: torch.Tensor, key: torch.Tensor) -> torch.Tensor: + """Local Q/K squared sums before TP all-reduce.""" + query = query.float() + key = key.float() + query_var = (query * query).sum(-1, keepdim=True) + key_var = (key * key).sum(-1, keepdim=True) + return torch.stack([query_var, key_var], dim=0) + + +def _qk_rmsnorm_apply(query: torch.Tensor, + key: torch.Tensor, + variance: torch.Tensor, + query_weight: torch.Tensor, + key_weight: torch.Tensor, + hidden_size: int, + eps: float): + """Apply whole-hidden Q/K RMSNorm from an already all-reduced variance.""" + dtype = query.dtype + query_var, key_var = variance / hidden_size + eps + query = (query.float() * torch.rsqrt(query_var)).to(dtype) * query_weight + key = (key.float() * torch.rsqrt(key_var)).to(dtype) * key_weight + return query, key + + +class Embedding(nn.Module): + """Token embedding container. + + Mirrors ``self.embedding`` in the reference: a plain ``nn.Module`` holding + a ``word_embeddings`` table. Keeping the attribute name ``word_embeddings`` + makes the checkpoint key ``embedding.word_embeddings.weight`` line up + directly with ``self.embedding.word_embeddings``. + """ + + def __init__(self, + config: PretrainedConfig, + dtype: torch.dtype = None, + device: torch.device = None): + super().__init__() + self.word_embeddings = build_embedding( + config.vocab_size, + config.hidden_size, + getattr(config, 'pad_token_id', None), + dtype=dtype, + device=device, + ) + + def forward(self, input_ids: torch.Tensor) -> torch.Tensor: + """forward.""" + return self.word_embeddings(input_ids) + + +class Quantizer(nn.Module): + """Rewrite of ``_Quantizer``. + + The reference keeps codebooks as a ``ParameterList`` only to produce + checkpoint keys ``codebook.0`` ... ``codebook.N``. Runtime computation only + needs the stacked tensor returned by ``transformed_codebook()``. + """ + + def __init__(self, + config: PretrainedConfig, + dtype: torch.dtype = None, + device: torch.device = None): + super().__init__() + self.num_codebooks = int(config.concept_v22_vq_num_codebooks) + self.codebook_size = int(config.concept_v22_vq_codebook_size) + hidden_size = int(config.hidden_size) + assert hidden_size % self.num_codebooks == 0, ( + f'hidden_size={hidden_size} must be divisible by num_codebooks={self.num_codebooks}') + self.codebook_dim = hidden_size // self.num_codebooks + self.hidden_size = hidden_size + self.codebook = nn.Parameter( + torch.empty( + self.num_codebooks, + self.codebook_size, + self.codebook_dim, + dtype=dtype, + device=device, + ), + requires_grad=False, + ) + + def transformed_codebook(self): + """Return codebook as ``[num_codebooks, codebook_size, + codebook_dim]``.""" + return self.codebook + + def forward(self, concept_logits: torch.Tensor) -> torch.Tensor: + """Convert per-codebook logits to hidden vectors. + + Args: + concept_logits: ``[..., num_codebooks, codebook_size]``. In the + LMDeploy engine the leading dims may be a packed continuous + batching dimension, e.g. ``[num_concepts_total]``. + + Returns: + Quantized/predicted vectors with shape ``[..., hidden_size]``. + """ + codebook = self.transformed_codebook().to(concept_logits.dtype) + vectors = torch.einsum('...hk,hkd->...hd', concept_logits, codebook) + return vectors.flatten(-2, -1) + + +class DepthDD(nn.Module): + """Rewrite of ``_DepthDD``. + + This is a small replicated per-token depth mixer. It does not mix tokens; + it computes ``num_prev`` route weights from the current hidden state and + combines the matching per-layer history states. The reference only handles + dense ``[seq, batch, hidden]`` tensors because it stacks history at + ``dim=2``. LMDeploy continuous batching commonly uses packed + ``[num_tokens, hidden]`` tensors. Runtime passes preallocated tensor + history to avoid repeated ``torch.stack`` copies. + """ + + def __init__(self, + config: PretrainedConfig, + layer_idx: int, + use_softmax: bool, + dtype: torch.dtype = None, + device: torch.device = None): + super().__init__() + self.num_prev = int(layer_idx) + 2 + route_hidden_size = self.num_prev + self.w1 = nn.Linear(config.hidden_size, route_hidden_size, bias=False, dtype=dtype, device=device) + self.w2 = nn.Linear(route_hidden_size, self.num_prev, bias=False, dtype=dtype, device=device) + self.static_a = nn.Parameter(torch.zeros(self.num_prev, dtype=dtype, device=device), requires_grad=False) + self.use_softmax = bool(use_softmax) + for param in self.parameters(): + param.requires_grad_(False) + + def _history_tensor(self, hidden_states: torch.Tensor, history_states: _HistoryStates, history_dim: int): + """Return history in shape ``[..., num_prev, hidden_size]``.""" + history_dim = history_dim if history_dim >= 0 else history_dim + history_states.dim() + if history_dim != history_states.dim() - 2: + history_states = history_states.movedim(history_dim, -2) + return history_states + + def forward(self, + hidden_states: torch.Tensor, + history_states: _HistoryStates, + history_dim: int = -2) -> torch.Tensor: + """forward.""" + history = self._history_tensor(hidden_states, history_states, history_dim) + weights = self.w2(F.gelu(self.w1(hidden_states))) + weights = weights + self.static_a.to(dtype=weights.dtype) + if self.use_softmax: + weights = weights.softmax(dim=-1) + return torch.einsum('...l,...lh->...h', weights, history) + + +class SelfDD(nn.Module): + """Rewrite of ``_SelfDD``.""" + + def __init__(self, + config: PretrainedConfig, + num_layers: int, + use_softmax: bool = False, + dtype: torch.dtype = None, + device: torch.device = None): + super().__init__() + self.num_layers = int(num_layers) + self.depth_dds = nn.ModuleList([ + DepthDD(config, layer_idx, use_softmax, dtype=dtype, device=device) + for layer_idx in range(self.num_layers) + ]) + + def make_history_buffer(self, hidden_states: torch.Tensor) -> torch.Tensor: + """Allocate layer-major history buffer ``[num_layers + 1, + *hidden_shape]``.""" + return hidden_states.new_empty((self.num_layers + 1, *hidden_states.shape)) + + @staticmethod + def write_history(history_buffer: torch.Tensor, slot_idx: int, hidden_states: torch.Tensor): + """Copy one history block into a layer-major history buffer.""" + history_buffer[int(slot_idx)].copy_(hidden_states) + return history_buffer + + @staticmethod + def history_view(history_buffer: torch.Tensor, layer_idx: int): + """Return layer-major history needed by ``layer_idx`` without copying. + + This is CUDA-graph safe when ``layer_idx`` is a Python constant for the + current layer and ``history_buffer`` has the fixed graph-capture shape. + """ + return history_buffer[:int(layer_idx) + 2] + + def forward_from_buffer(self, + layer_idx: int, + hidden_states: torch.Tensor, + history_buffer: torch.Tensor): + """Runtime path: fixed buffer, no list materialization.""" + layer_idx = int(layer_idx) + return self.depth_dds[layer_idx]( + hidden_states, + self.history_view(history_buffer, layer_idx), + history_dim=0, + ) + + def forward(self, + layer_idx: int, + hidden_states: torch.Tensor, + history_states: _HistoryStates): + """forward.""" + return self.depth_dds[int(layer_idx)](hidden_states, history_states) + + +class ResidualRoute(nn.Module): + """Rewrite of ``_ResidualRoute``. + + This module computes a source-state mixture and adds it as a gated residual update to the target hidden state. It is + small and replicated. + """ + + def __init__(self, + config: PretrainedConfig, + num_source_states: int, + use_softmax: bool = True, + dtype: torch.dtype = None, + device: torch.device = None): + super().__init__() + self.num_source_states = int(num_source_states) + route_hidden_size = max(1, self.num_source_states) + self.w1 = nn.Linear(config.hidden_size, route_hidden_size, bias=False, dtype=dtype, device=device) + self.w2 = nn.Linear(route_hidden_size, self.num_source_states, bias=False, dtype=dtype, device=device) + self.residual_diag = nn.Parameter(torch.zeros(config.hidden_size, dtype=dtype, device=device), + requires_grad=False) + self.use_softmax = bool(use_softmax) + for param in self.parameters(): + param.requires_grad_(False) + + def _source_tensor(self, + target_hidden: torch.Tensor, + source_states: _SourceStates, + source_dim: int): + """Return source states in shape ``[..., active_sources, + hidden_size]``.""" + if source_states is None: + return None + + source_dim = source_dim if source_dim >= 0 else source_dim + source_states.dim() + if source_states.shape[source_dim] == 0: + return None + if source_dim != source_states.dim() - 2: + source_states = source_states.movedim(source_dim, -2) + return source_states + + def _route_weights(self, target_hidden: torch.Tensor, active_sources: int): + """Compute route weights and keep the last active source logits.""" + weights = self.w2(F.gelu(self.w1(target_hidden))) + weights = weights[..., -active_sources:] + if self.use_softmax: + weights = weights.softmax(dim=-1) + return weights + + def _add_update(self, + target_hidden: torch.Tensor, + source_mix: torch.Tensor, + residual_scale: torch.Tensor | None = None): + """Apply residual diagonal and optional scale, then add to target.""" + update = source_mix * self.residual_diag.to(source_mix.dtype) + if residual_scale is not None: + update = update * residual_scale.to(update.dtype) + return target_hidden + update.to(target_hidden.dtype) + + def forward(self, + target_hidden: torch.Tensor, + source_states: _SourceStates, + residual_scale: torch.Tensor | None = None, + source_dim: int = -2): + """Mix source states into ``target_hidden``.""" + source_states = self._source_tensor(target_hidden, source_states, source_dim) + if source_states is None: + return target_hidden + active_sources = source_states.shape[-2] + weights = self._route_weights(target_hidden, active_sources) + source_mix = torch.einsum('...m,...mh->...h', weights, source_states) + return self._add_update(target_hidden, source_mix, residual_scale) + + +class ConceptRoute(nn.Module): + """Rewrite of ``_ConceptRoute``. + + Applies LayerNorm to the final concept state, scales it elementwise with a learned diagonal, optionally applies a + route scale, then adds the update to decoder hidden states. This is replicated and token-local. + """ + + def __init__(self, + config: PretrainedConfig, + dtype: torch.dtype = None, + device: torch.device = None): + super().__init__() + self.concept_norm = nn.LayerNorm(config.hidden_size, + eps=getattr(config, 'layernorm_epsilon', 1e-6), + dtype=dtype, + device=device) + self.final_diag = nn.Parameter(torch.zeros(config.hidden_size, dtype=dtype, device=device), + requires_grad=False) + for param in self.parameters(): + param.requires_grad_(False) + + def forward(self, + hidden_states: torch.Tensor, + final_concept_state: torch.Tensor, + final_scale: torch.Tensor | None = None): + """forward.""" + concept = self.concept_norm(final_concept_state) + update = concept * self.final_diag.to(concept.dtype) + if final_scale is not None: + update = update * final_scale.to(update.dtype) + return hidden_states + update.to(hidden_states.dtype) + + +class TwoRouteAdd(nn.Module): + """Rewrite of ``_TwoRouteAdd``. + + It first applies decoder-side ``_DepthDD`` to the decoder history, then + injects the final concept state with ``_ConceptRoute``. + """ + + def __init__(self, + config: PretrainedConfig, + dtype: torch.dtype = None, + device: torch.device = None): + super().__init__() + self.num_layers = int(config.concept_decoder_layers) + use_softmax = bool(getattr(config, 'concept_dd_two_route_add_decoder_use_softmax', True)) + self.decoder_dds = nn.ModuleList([ + DepthDD(config, layer_idx, use_softmax, dtype=dtype, device=device) + for layer_idx in range(self.num_layers) + ]) + self.concept_routes = nn.ModuleList([ + ConceptRoute(config, dtype=dtype, device=device) + for _ in range(self.num_layers) + ]) + + def make_history_buffer(self, hidden_states: torch.Tensor) -> torch.Tensor: + """Allocate layer-major decoder history buffer ``[num_layers + 1, + *hidden_shape]``.""" + return hidden_states.new_empty((self.num_layers + 1, *hidden_states.shape)) + + @staticmethod + def write_history(history_buffer: torch.Tensor, slot_idx: int, hidden_states: torch.Tensor): + """Copy one decoder history block into a layer-major history buffer.""" + return SelfDD.write_history(history_buffer, slot_idx, hidden_states) + + @staticmethod + def history_view(history_buffer: torch.Tensor, layer_idx: int): + """Return layer-major decoder history needed by ``layer_idx`` without + copying.""" + return SelfDD.history_view(history_buffer, layer_idx) + + def forward_from_buffer(self, + layer_idx: int, + hidden_states: torch.Tensor, + history_buffer: torch.Tensor, + final_concept_state: torch.Tensor, + final_scale: torch.Tensor | None = None): + """Runtime path: fixed decoder history buffer, no list materialization.""" + layer_idx = int(layer_idx) + hidden_states = self.decoder_dds[layer_idx]( + hidden_states, + self.history_view(history_buffer, layer_idx), + history_dim=0, + ) + return self.concept_routes[layer_idx](hidden_states, final_concept_state, final_scale) + + +class PredictionHeads(nn.Module): + """Merged per-codebook prediction heads. + + Native checkpoints store one ``prediction_heads.N`` linear per concept + codebook. Runtime uses one merged projection and loads each native head into + a deterministic output shard through LMDeploy's ``param.weight_loader`` + pattern. + """ + + def __init__(self, + config: PretrainedConfig, + dtype: torch.dtype = None, + device: torch.device = None, + prefix: str = ''): + codebook_size = int(config.concept_v22_vq_codebook_size) + num_codebooks = int(config.concept_v22_vq_num_codebooks) + super().__init__() + self.num_codebooks = num_codebooks + self.codebook_size = codebook_size + quantization_config = getattr(config, 'quantization_config', None) + self.proj = build_merged_colwise_linear( + config.hidden_size, + [codebook_size] * num_codebooks, + bias=True, + dtype=dtype, + device=device, + quant_config=quantization_config, + # Keep concept logits replicated for now. Output TP would make the + # trailing ``[num_codebooks, codebook_size]`` contract sharded and + # needs a defined distributed sampling/gather path first. + is_tp=False, + out_names=list(range(num_codebooks)), + prefix=add_prefix('proj', prefix), + ) + + def forward(self, hidden_states: torch.Tensor): + """Return logits in shape ``[..., num_codebooks, codebook_size]``.""" + logits = self.proj(hidden_states) + return logits.unflatten(-1, (self.num_codebooks, self.codebook_size)) + + +class ConceptPredictor(nn.Module): + """Rewrite of ``_ConceptPredictor``. + + The predictor owns the high-level concept OLMo block, per-codebook prediction heads, concept self-DD, encoder-read + routes, and the shared source LayerNorm used before encoder states are routed into the concept stream. + """ + + def __init__(self, + config: PretrainedConfig, + dtype: torch.dtype = None, + device: torch.device = None, + prefix: str = ''): + super().__init__() + self.num_layers = int(config.concept_special_layers) + self.num_codebooks = int(config.concept_v22_vq_num_codebooks) + self.codebook_size = int(config.concept_v22_vq_codebook_size) + self.hlm_block = OlmoBlock(config, + self.num_layers, + post_layer_norm=True, + dtype=dtype, + device=device, + prefix=add_prefix('hlm_block', prefix)) + self.prediction_heads = PredictionHeads(config, + dtype=dtype, + device=device, + prefix=add_prefix('prediction_heads', prefix)) + self.concept_self_dd = SelfDD(config, + self.num_layers, + use_softmax=False, + dtype=dtype, + device=device) + self.concept_read_encoder_routes = nn.ModuleList([ + ResidualRoute(config, + int(config.concept_encoder_layers) - 1, + use_softmax=True, + dtype=dtype, + device=device) + for _ in range(self.num_layers) + ]) + self.concept_read_encoder_shared_source_norm = nn.LayerNorm( + config.hidden_size, + eps=getattr(config, 'layernorm_epsilon', 1e-6), + dtype=dtype, + device=device) + for param in self.concept_read_encoder_shared_source_norm.parameters(): + param.requires_grad_(False) + + def normalize_encoder_concept_states(self, encoder_concept_states: torch.Tensor): + """Apply the shared source norm used before concept-read-encoder + routes.""" + return self.concept_read_encoder_shared_source_norm(encoder_concept_states) + + def forward(self, + concept_hidden: torch.Tensor, + encoder_concept_states: torch.Tensor, + position_ids: torch.Tensor, + past_key_values: list[list[torch.Tensor]] | None = None, + attn_metadata: Any = None, + encoder_source_dim: int = -2): + """Concept predictor forward. + + This mirrors the reference control flow, but uses LMDeploy's OLMo layer + rewrite and paged attention inputs. Full end-to-end use requires the + caller to provide concept-stream ``past_key_values`` and + ``attn_metadata`` matching the concept token layout. + """ + hidden_states = concept_hidden + history_buffer = self.concept_self_dd.make_history_buffer(hidden_states) + self.concept_self_dd.write_history(history_buffer, 0, hidden_states) + raw_states = [] + rotary_pos_emb = self.hlm_block._make_rotary_pos_emb(hidden_states, position_ids) + + for layer_idx, layer in enumerate(self.hlm_block.layers): + pkv = past_key_values[layer_idx] if past_key_values is not None else None + raw = layer( + hidden_states, + rotary_pos_emb=rotary_pos_emb, + past_key_value=pkv, + attn_metadata=attn_metadata, + ) + raw_states.append(raw) + self.concept_self_dd.write_history(history_buffer, layer_idx + 1, raw) + hidden_states = self.concept_self_dd.forward_from_buffer(layer_idx, raw, history_buffer) + hidden_states = self.concept_read_encoder_routes[layer_idx]( + hidden_states, + encoder_concept_states, + source_dim=encoder_source_dim, + ) + + hidden_states = self.hlm_block.final_layernorm(hidden_states) + logits = self.prediction_heads(hidden_states) + return logits, raw_states + + +class OlmoAttention(nn.Module): + """Rewrite of ``_OlmoSelfAttention``. + + Differences from the reference: + - fused QKV uses lmdeploy's ``build_qkv_proj`` (standard [Q,K,V] packing, + TP-aware). The reference stores QKV per attention head, so the native + checkpoint weight is re-laid-out during ``load_weights``. + - q/k RMSNorm operates on the whole hidden_size (as in the reference), + not per-head. TP>1 needs an all-reduced variance; we follow internvl's + qkv_norm pattern. + - attention runs through lmdeploy's paged ``Attention`` primitive. + """ + + def __init__(self, + config: PretrainedConfig, + dtype: torch.dtype = None, + device: torch.device = None, + sliding_window: int | None = None, + prefix: str = ''): + super().__init__() + quantization_config = getattr(config, 'quantization_config', None) + hidden_size = config.hidden_size + num_heads = config.num_attention_heads + head_dim = config.kv_channels + num_kv_heads = num_heads + assert hidden_size == num_heads * head_dim, ( + f'ConceptLM OLMo attention expects hidden_size == num_heads * head_dim, ' + f'got {hidden_size} != {num_heads} * {head_dim}.') + + self.num_heads = num_heads + self.num_kv_heads = num_kv_heads + self.head_dim = head_dim + self.hidden_size = hidden_size + self.layernorm_epsilon = config.layernorm_epsilon + + # packed qkv + self.qkv_proj = build_qkv_proj( + hidden_size, + num_q_heads=num_heads, + num_kv_heads=num_kv_heads, + head_size=head_dim, + bias=False, + quant_config=quantization_config, + dtype=dtype, + device=device, + prefix=add_prefix('qkv_proj', prefix), + ) + + # q, k norm over the whole hidden_size (NOT per-head). tp=True with + # head_dim alignment so the weight shards correctly under TP. + self.q_layernorm = RMSNorm(hidden_size, + config.layernorm_epsilon, + quant_config=quantization_config, + dtype=dtype, + device=device, + tp=True, + align=head_dim, + prefix=add_prefix('q_layernorm', prefix)) + self.k_layernorm = RMSNorm(hidden_size, + config.layernorm_epsilon, + quant_config=quantization_config, + dtype=dtype, + device=device, + tp=True, + align=head_dim, + prefix=add_prefix('k_layernorm', prefix)) + + # rotary embedding + self.apply_rotary_pos_emb = ApplyRotaryEmb() + + # attention + self.attn_fwd = Attention(num_heads, + head_dim, + num_kv_heads=num_kv_heads, + v_head_size=head_dim, + sliding_window=None if sliding_window is None else int(sliding_window), + device=device) + + # o_proj + self.o_proj = build_o_proj(num_heads * head_dim, + hidden_size, + bias=False, + quant_config=quantization_config, + dtype=dtype, + device=device, + is_tp=True, + prefix=add_prefix('o_proj', prefix)) + + def _qkv_norm(self, q: torch.Tensor, k: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """RMSNorm over the whole hidden_size, TP-correct via all-reduce.""" + import lmdeploy.pytorch.distributed as dist + q_shape = q.shape + k_shape = k.shape + q = q.flatten(-2, -1) + k = k.flatten(-2, -1) + + tp, _ = dist.get_tp_world_rank('attn') + if tp == 1: + q = self.q_layernorm(q) + k = self.k_layernorm(k) + return q.view(q_shape), k.view(k_shape) + + # variance is computed over the full hidden_size, so it must be + # all-reduced across TP ranks before normalizing the local shard. + variance = _qk_rmsnorm_variance(q, k) + dist.all_reduce(variance) + q, k = _qk_rmsnorm_apply(q, k, variance, self.q_layernorm.weight, + self.k_layernorm.weight, self.hidden_size, + self.layernorm_epsilon) + return q.view(q_shape), k.view(k_shape) + + def forward(self, + hidden_states: torch.Tensor, + rotary_pos_emb: tuple[torch.Tensor, torch.Tensor], + past_key_value: list[torch.Tensor] | None = None, + attn_metadata: Any = None): + """Rewrite of _OlmoSelfAttention.forward.""" + # qkv proj -> (batch, seq, num_heads, head_dim) each + qkv_states = self.qkv_proj(hidden_states) + qkv_states = qkv_states.flatten(0, -2) # (-1, heads_total, head_dim) + query_states, key_states, value_states = self.qkv_proj.split_qkv(qkv_states) + + # q, k norm (whole hidden_size) + query_states, key_states = self._qkv_norm(query_states, key_states) + + # rotary embedding + cos, sin = rotary_pos_emb + query_states, key_states = self.apply_rotary_pos_emb( + query_states, key_states, cos, sin, inplace=True) + + # attention (paged) + attn_output = self.attn_fwd( + query_states, + key_states, + value_states, + past_key_value[0], + past_key_value[1], + attn_metadata, + k_scales_zeros=None if len(past_key_value) == 2 else past_key_value[2], + v_scales_zeros=None if len(past_key_value) == 2 else past_key_value[3], + inplace=True, + ) + attn_output = attn_output.reshape(*hidden_states.shape[:-1], -1) + + # o proj + attn_output = self.o_proj(attn_output) + return attn_output + + +class OlmoMLP(nn.Module): + """Rewrite of ``_OlmoMLP`` (SwiGLU).""" + + def __init__(self, + config: PretrainedConfig, + dtype: torch.dtype = None, + device: torch.device = None, + prefix: str = ''): + super().__init__() + quantization_config = getattr(config, 'quantization_config', None) + self.gate_up_proj = build_gateup_linear( + config.hidden_size, + [config.ffn_hidden_size, config.ffn_hidden_size], + bias=False, + dtype=dtype, + device=device, + quant_config=quantization_config, + is_tp=True, + prefix=add_prefix('gate_up_proj', prefix), + ) + self.act_fn = SiluAndMul(inplace=True) + self.down_proj = build_down_linear( + config.ffn_hidden_size, + config.hidden_size, + bias=False, + quant_config=quantization_config, + dtype=dtype, + device=device, + is_tp=True, + prefix=add_prefix('down_proj', prefix), + ) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + """forward.""" + gate_up = self.gate_up_proj(hidden_states) + act = self.act_fn(gate_up) + return self.down_proj(act) + + +class OlmoLayer(nn.Module): + """Rewrite of ``_OlmoLayer``. + + The reference uses post-norm residuals:: + h = h + post_attention_layernorm(attn(h)) + h = h + post_feedforward_layernorm(mlp(h)) + This block therefore calls RMSNorm without the residual argument and adds + the residual explicitly. + """ + + def __init__(self, + config: PretrainedConfig, + layer_idx: int, + dtype: torch.dtype = None, + device: torch.device = None, + sliding_window: int | None = None, + prefix: str = ''): + super().__init__() + quantization_config = getattr(config, 'quantization_config', None) + self.layer_number = int(layer_idx) + 1 + self.self_attention = OlmoAttention(config, + dtype=dtype, + device=device, + sliding_window=sliding_window, + prefix=add_prefix('self_attention', prefix)) + self.post_attention_layernorm = RMSNorm(config.hidden_size, + config.layernorm_epsilon, + quant_config=quantization_config, + dtype=dtype, + device=device, + prefix=add_prefix('post_attention_layernorm', prefix)) + self.mlp = OlmoMLP(config, dtype=dtype, device=device, prefix=add_prefix('mlp', prefix)) + self.post_feedforward_layernorm = RMSNorm(config.hidden_size, + config.layernorm_epsilon, + quant_config=quantization_config, + dtype=dtype, + device=device, + prefix=add_prefix('post_feedforward_layernorm', prefix)) + + def forward(self, + hidden_states: torch.Tensor, + rotary_pos_emb: tuple[torch.Tensor, torch.Tensor], + past_key_value: list[torch.Tensor] | None = None, + attn_metadata: Any = None): + """forward. + + The reference uses post-norm residuals:: + h = h + post_attention_layernorm(attn(h)) + h = h + post_feedforward_layernorm(mlp(h)) + This is NOT the same as lmdeploy's pre-norm residual form + ``norm(x, residual=h)`` (which normalizes x+h). So we call the norm + without a residual and add explicitly. + """ + attn_out = self.self_attention( + hidden_states=hidden_states, + rotary_pos_emb=rotary_pos_emb, + past_key_value=past_key_value, + attn_metadata=attn_metadata, + ) + hidden_states = hidden_states + self.post_attention_layernorm(attn_out) + + mlp_out = self.mlp(hidden_states) + hidden_states = hidden_states + self.post_feedforward_layernorm(mlp_out) + return hidden_states + + +class OlmoBlock(nn.Module): + """Rewrite of ``_OlmoBlock``. + + Holds a stack of ``_OlmoLayer`` and an optional final RMSNorm. Mirrors the + reference's ``forward`` return contract: ``(hidden_states, layer_states)``. + """ + + def __init__(self, + config: PretrainedConfig, + num_layers: int, + post_layer_norm: bool, + dtype: torch.dtype = None, + device: torch.device = None, + window_size: Any = _CONFIG_VALUE, + skip_frequency: Any = _CONFIG_VALUE, + prefix: str = ''): + super().__init__() + quantization_config = getattr(config, 'quantization_config', None) + if window_size is _CONFIG_VALUE: + window_size = _get_configured_window(config) + elif isinstance(window_size, (list, tuple)): + window_size = window_size[0] if len(window_size) > 0 else None + if window_size is not None: + window_size = int(window_size) + if skip_frequency is _CONFIG_VALUE: + skip_frequency = getattr(config, 'window_attn_skip_freq', None) + if skip_frequency is not None: + skip_frequency = int(skip_frequency) + + self.layers = nn.ModuleList([ + OlmoLayer(config, + layer_idx, + dtype=dtype, + device=device, + sliding_window=self._layer_sliding_window(layer_idx + 1, window_size, skip_frequency), + prefix=add_prefix(f'layers.{layer_idx}', prefix)) + for layer_idx in range(num_layers) + ]) + self.final_layernorm = ( + RMSNorm(config.hidden_size, + config.layernorm_epsilon, + quant_config=quantization_config, + dtype=dtype, + device=device, + prefix=add_prefix('final_layernorm', prefix)) if post_layer_norm else None) + self.rotary_emb = _make_olmo_rotary_embedding(config, device=device) + self.num_heads = int(config.num_attention_heads) + self.head_dim = int(config.kv_channels) + + @staticmethod + def _layer_sliding_window(layer_number: int, window_size: int | None, skip_frequency: int | None): + """Match reference OLMo window/full attention alternation.""" + if window_size is None: + return None + if skip_frequency is not None and layer_number % skip_frequency == 0: + return None + return window_size + + def _make_rotary_pos_emb(self, hidden_states: torch.Tensor, position_ids: torch.Tensor): + """Create RoPE cos/sin in the shape consumed by ``ApplyRotaryEmb``.""" + if position_ids.dim() == 1: + position_ids = position_ids.unsqueeze(0) + cos, sin = self.rotary_emb(hidden_states, position_ids) + if cos.dim() == 3 and cos.size(0) == 1: + cos = cos[0] + sin = sin[0] + return cos, sin + + def forward(self, + hidden_states: torch.Tensor, + position_ids: torch.Tensor, + past_key_values: list[list[torch.Tensor]] | None = None, + attn_metadata: Any = None, + collect: bool = False): + """forward.""" + layer_states = [] + rotary_pos_emb = self._make_rotary_pos_emb(hidden_states, position_ids) + for idx, layer in enumerate(self.layers): + pkv = past_key_values[idx] if past_key_values is not None else None + hidden_states = layer( + hidden_states, + rotary_pos_emb=rotary_pos_emb, + past_key_value=pkv, + attn_metadata=attn_metadata, + ) + if collect: + layer_states.append(hidden_states) + if self.final_layernorm is not None: + hidden_states = self.final_layernorm(hidden_states) + return hidden_states, layer_states + + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]], prefix: str = ''): + """Load native ConceptLM OLMo block weights into the LMDeploy + rewrite.""" + if prefix and not prefix.endswith('.'): + prefix = f'{prefix}.' + + params_dict = dict(self.named_parameters()) + loaded_names = set() + for name, loaded_weight in weights: + if 'rotary_emb.inv_freq' in name: + continue + if prefix: + if not name.startswith(prefix): + continue + name = name[len(prefix):] + + if name.endswith('.self_attention.qkv_proj.weight'): + param = params_dict[name] + query, key, value = param.weight_spliter(loaded_weight) + load_weight(param, query, shard_id='q') + load_weight(param, key, shard_id='k') + load_weight(param, value, shard_id='v') + loaded_names.add(name) + continue + + if name.endswith('.self_attention.linear_qkv.weight'): + target_name = name.replace('.self_attention.linear_qkv.weight', + '.self_attention.qkv_proj.weight') + param = params_dict[target_name] + loaded_weight = _repack_olmo_qkv_weight(loaded_weight, self.num_heads, self.head_dim) + query, key, value = param.weight_spliter(loaded_weight) + load_weight(param, query, shard_id='q') + load_weight(param, key, shard_id='k') + load_weight(param, value, shard_id='v') + loaded_names.add(target_name) + continue + + if name.endswith('.self_attention.linear_proj.weight'): + target_name = name.replace('.self_attention.linear_proj.weight', + '.self_attention.o_proj.weight') + elif name.endswith('.mlp.gate_up_proj.weight'): + param = params_dict[name] + gate, up = param.weight_spliter(loaded_weight) + load_weight(param, gate, shard_id=0) + load_weight(param, up, shard_id=1) + loaded_names.add(name) + continue + elif name.endswith('.mlp.linear_fc1.weight'): + target_name = name.replace('.mlp.linear_fc1.weight', '.mlp.gate_up_proj.weight') + param = params_dict[target_name] + gate, up = param.weight_spliter(loaded_weight) + load_weight(param, gate, shard_id=0) + load_weight(param, up, shard_id=1) + loaded_names.add(target_name) + continue + elif name.endswith('.mlp.linear_fc2.weight'): + target_name = name.replace('.mlp.linear_fc2.weight', '.mlp.down_proj.weight') + else: + target_name = name + + param = params_dict.get(target_name) + if param is None: + continue + load_weight(param, loaded_weight) + loaded_names.add(target_name) + return loaded_names diff --git a/lmdeploy/pytorch/models/intern_ncp/weight.py b/lmdeploy/pytorch/models/intern_ncp/weight.py new file mode 100644 index 0000000000..aee434dc15 --- /dev/null +++ b/lmdeploy/pytorch/models/intern_ncp/weight.py @@ -0,0 +1,23 @@ +# Copyright (c) OpenMMLab. All rights reserved. +import torch + + +def _repack_olmo_qkv_weight(loaded_weight: torch.Tensor, num_heads: int, head_dim: int): + """Convert native OLMo per-head [Q,K,V] QKV packing to LMDeploy + [Q][K][V].""" + leading_shape = loaded_weight.shape[1:] + loaded_weight = loaded_weight.reshape(num_heads, 3, head_dim, *leading_shape) + query = loaded_weight[:, 0].flatten(0, 1) + key = loaded_weight[:, 1].flatten(0, 1) + value = loaded_weight[:, 2].flatten(0, 1) + return torch.cat([query, key, value], dim=0) + + +def _load_stacked_codebook_weight(param: torch.nn.Parameter, loaded_weight: torch.Tensor, codebook_idx: int): + """Load one native ``codebook.N`` checkpoint tensor into a stacked + codebook.""" + assert 0 <= codebook_idx < param.size(0), f'Invalid codebook index: {codebook_idx}' + target = param.data[codebook_idx] + assert target.size() == loaded_weight.size(), ( + f'Attempted to load codebook weight ({loaded_weight.size()}) into parameter slice ({target.size()})') + target.copy_(loaded_weight) diff --git a/lmdeploy/pytorch/nn/conceptlm.py b/lmdeploy/pytorch/nn/conceptlm.py index 22f12bcf31..d36e037737 100644 --- a/lmdeploy/pytorch/nn/conceptlm.py +++ b/lmdeploy/pytorch/nn/conceptlm.py @@ -25,9 +25,8 @@ def decode_chunk_state_update( position_ids: Tensor, chunk_size: int, merge_method: str, - ) -> tuple[Tensor, Tensor, Tensor]: - """Update state cache and return concept inputs, next rows, and - mask.""" + ) -> tuple[Tensor, Tensor]: + """Update state cache and return concept inputs plus update mask.""" return self.impl.decode_chunk_state_update( chunk_source_state_cache, current_source_states, @@ -91,22 +90,3 @@ def decode_concept_state_update( state_ids, update_mask, ) - - def forward( - self, - chunk_source_state_cache: Tensor, - current_source_states: Tensor, - state_ids: Tensor, - position_ids: Tensor, - chunk_size: int, - merge_method: str, - ) -> tuple[Tensor, Tensor, Tensor]: - """Alias the current runtime op for module-call compatibility.""" - return self.decode_chunk_state_update( - chunk_source_state_cache, - current_source_states, - state_ids, - position_ids, - chunk_size, - merge_method, - ) From 8181d9d3a7905478c1ad8e10ab5f48a63ec6e250 Mon Sep 17 00:00:00 2001 From: grimoire Date: Sun, 26 Jul 2026 16:26:00 +0800 Subject: [PATCH 07/16] Move ConceptLM runtime layout ops to backend --- lmdeploy/pytorch/backends/conceptlm.py | 595 ++++++++++++++++- lmdeploy/pytorch/backends/cuda/conceptlm.py | 4 +- .../pytorch/backends/default/conceptlm.py | 20 +- .../pytorch/models/intern_ncp/metadata.py | 87 --- .../pytorch/models/intern_ncp/modeling.py | 607 ++---------------- lmdeploy/pytorch/nn/conceptlm.py | 126 +++- 6 files changed, 765 insertions(+), 674 deletions(-) diff --git a/lmdeploy/pytorch/backends/conceptlm.py b/lmdeploy/pytorch/backends/conceptlm.py index f6c279c724..5a819e6eec 100644 --- a/lmdeploy/pytorch/backends/conceptlm.py +++ b/lmdeploy/pytorch/backends/conceptlm.py @@ -1,7 +1,107 @@ # Copyright (c) OpenMMLab. All rights reserved. from abc import ABC, abstractmethod +from dataclasses import dataclass, replace +from typing import Any +import torch from torch import Tensor +from torch.nn import functional as F +from transformers.configuration_utils import PretrainedConfig + + +@dataclass +class ConceptChunkStateUpdateResult: + """Fixed-shape result of one decode chunk-source state update.""" + + concept_input_states: Tensor + concept_update_mask: Tensor + + +@dataclass +class ConceptDecodeMetadata: + """Fixed-layout decode metadata derived once from engine inputs.""" + + position_ids: Tensor + state_ids: Tensor + safe_state_ids: Tensor + valid_state_mask: Tensor + + +@dataclass +class ConceptPrefillMetadata: + """Packed ConceptLM prefill runtime metadata. + + The model treats this as a backend-owned plan. Fields stay public for the current torch fallback path and tests, but + model code should not rebuild or reinterpret this layout directly. + """ + + token_q_seqlens: Tensor + token_q_start_loc: Tensor + concept_q_seqlens: Tensor + concept_q_start_loc: Tensor + concept_position_ids: Tensor + merge_token_to_concept: Tensor + merge_token_counts: Tensor + merge_first_token_ids: Tensor + merge_last_token_ids: Tensor + merge_short_concept_mask: Tensor + token_to_concept: Tensor + num_tokens_total: int + num_concepts_total: int + max_concepts_per_request: int + + +@dataclass +class _PrefillTokenLayout: + """Packed token-stream layout derived from prefill attention metadata.""" + + q_seqlens: Tensor + q_start_loc: Tensor + q_seqlens_long: Tensor + q_start_loc_long: Tensor + token_seq: Tensor + token_pos: Tensor + total_tokens: int + + +@dataclass +class _PrefillConceptLayout: + """Compact chunk-token stream layout used by concept predictor prefill.""" + + q_seqlens: Tensor + q_seqlens_long: Tensor + q_start_loc: Tensor + q_start_loc_long: Tensor + seq: Tensor + local_ids: Tensor + position_ids: Tensor + num_total: int + max_per_request: int + + +@dataclass +class _PrefillMergeLayout: + """Token-to-concept merge metadata for compact prefill.""" + + token_to_concept: Tensor + token_counts: Tensor + first_token_ids: Tensor + last_token_ids: Tensor + short_concept_mask: Tensor + + +def _flatten_decode_position_ids(position_ids: Tensor, batch_size: int, device: torch.device) -> Tensor: + """Normalize decode position ids to one absolute position per batch row.""" + if position_ids.dim() == 0: + position_ids = position_ids.view(1) + if position_ids.dim() == 1: + return position_ids.to(device=device, dtype=torch.long) + position_ids = position_ids.reshape(-1) + if position_ids.numel() == batch_size: + return position_ids.to(device=device, dtype=torch.long) + assert position_ids.numel() % batch_size == 0, ( + f'Cannot map position_ids with {position_ids.numel()} elements to batch size {batch_size}.') + return position_ids.reshape(-1, batch_size)[-1].to(device=device, dtype=torch.long) class ConceptLMRuntimeOpsImpl(ABC): @@ -11,6 +111,499 @@ class ConceptLMRuntimeOpsImpl(ABC): direct kernel calls while avoiding one OpType/nn module per small ConceptLM state operation. """ + def __init__(self, config: PretrainedConfig): + self.config = config + self.chunk_size = int(config.concept_chunk_size) + self.merge_method = getattr(config, 'concept_chunk_merge_method', 'meanpooling') + self.shift_feature = bool(getattr(config, 'concept_shift_feature', True)) + + @staticmethod + def concept_count_from_seq_len(seq_len: int, chunk_size: int) -> int: + """Return reference ConceptLM chunk count for one request length.""" + seq_len = int(seq_len or 0) + if seq_len <= 0: + return 0 + if seq_len < chunk_size: + return 1 + return seq_len // chunk_size + + @staticmethod + def concept_counts_from_q_seqlens(q_seqlens: Tensor, chunk_size: int) -> Tensor: + """Vectorized version of ``concept_count_from_seq_len``.""" + counts = torch.div(q_seqlens, chunk_size, rounding_mode='floor').clamp(min=1) + return torch.where(q_seqlens > 0, counts, torch.zeros_like(q_seqlens)) + + @staticmethod + def repeat_slot_ids(token_pos: Tensor, chunk_size: int, shift_feature: bool) -> Tensor: + """Return local concept slot read by each token after shift + semantics.""" + if shift_feature: + return torch.div(token_pos + 1, chunk_size, rounding_mode='floor') - 1 + return torch.div(token_pos, chunk_size, rounding_mode='floor') - 1 + + def flatten_decode_position_ids(self, position_ids: Tensor, batch_size: int, device: torch.device) -> Tensor: + """Normalize decode position ids to one absolute position per batch + row.""" + return _flatten_decode_position_ids(position_ids, batch_size, device) + + def build_decode_metadata(self, + position_ids: Tensor, + state_ids: Tensor | None, + batch_size: int, + device: torch.device) -> ConceptDecodeMetadata: + """Build fixed-shape decode metadata from engine state ids.""" + if state_ids is None: + raise RuntimeError('ConceptLM decode requires state_ids.') + state_ids = state_ids.to(device=device, dtype=torch.long).reshape(-1) + if state_ids.numel() != batch_size: + raise ValueError(f'Expected {batch_size} decode state ids, got {state_ids.numel()}.') + valid_state_mask = state_ids >= 0 + return ConceptDecodeMetadata( + position_ids=position_ids, + state_ids=state_ids, + safe_state_ids=state_ids.clamp(min=0), + valid_state_mask=valid_state_mask, + ) + + @staticmethod + def select_decode_state_rows(state_cache: Tensor, decode_metadata: ConceptDecodeMetadata) -> Tensor: + """Gather state-cache rows and zero out padded decode rows.""" + rows = state_cache.index_select(0, decode_metadata.safe_state_ids) + mask_shape = (decode_metadata.valid_state_mask.size(0), ) + (1, ) * (rows.dim() - 1) + valid_mask = decode_metadata.valid_state_mask.view(mask_shape) + return torch.where(valid_mask, rows, torch.zeros_like(rows)) + + def select_decode_last_state_rows(self, + last_state: Tensor | None, + last_final_state: Tensor, + last_raw_states: Tensor, + decode_metadata: ConceptDecodeMetadata) -> tuple[Tensor, Tensor]: + """Gather packed last-concept state rows once and return final/raw + views.""" + if last_state is not None: + rows = self.select_decode_state_rows(last_state, decode_metadata) + return rows[:, 0], rows[:, 1:] + + return ( + self.select_decode_state_rows(last_final_state, decode_metadata), + self.select_decode_state_rows(last_raw_states, decode_metadata), + ) + + def decode_concept_read_mask(self, decode_metadata: ConceptDecodeMetadata) -> Tensor: + """Return rows whose current decode token should read a cached + concept.""" + repeat_slots = self.repeat_slot_ids( + decode_metadata.position_ids, + self.chunk_size, + self.shift_feature, + ) + return decode_metadata.valid_state_mask & (repeat_slots >= 0) + + def build_concept_decode_metadata_static(self, token_attn_metadata: Any, + decode_metadata: ConceptDecodeMetadata): + """Build fixed-shape concept-stream decode metadata.""" + device = decode_metadata.position_ids.device + batch_size = decode_metadata.position_ids.numel() + q_seqlens = token_attn_metadata.q_seqlens + q_start_loc = token_attn_metadata.q_start_loc + kv_seqlens = token_attn_metadata.kv_seqlens + q_dtype = q_seqlens.dtype + q_start_dtype = q_start_loc.dtype + kv_dtype = kv_seqlens.dtype + + concept_q_seqlens = torch.ones((batch_size, ), dtype=q_dtype, device=device) + concept_q_start_loc = torch.arange(batch_size, dtype=q_start_dtype, device=device) + concept_cu_seqlens = F.pad(torch.cumsum(concept_q_seqlens, dim=0, dtype=torch.int32), (1, 0)) + concept_kv_seqlens = torch.div( + decode_metadata.position_ids + 1, + self.chunk_size, + rounding_mode='floor', + ).clamp(min=1).to(dtype=kv_dtype) + + updates = dict( + is_decoding=True, + block_offsets=token_attn_metadata.block_offsets, + q_start_loc=concept_q_start_loc, + q_seqlens=concept_q_seqlens, + kv_seqlens=concept_kv_seqlens, + cu_seqlens_q=concept_cu_seqlens, + cu_seqlens_k=concept_cu_seqlens, + ) + if hasattr(token_attn_metadata, 'kv_start_loc'): + updates['kv_start_loc'] = concept_kv_seqlens - concept_q_seqlens.to(dtype=concept_kv_seqlens.dtype) + if hasattr(token_attn_metadata, 'kv_flatten_size'): + updates['kv_flatten_size'] = batch_size + if hasattr(token_attn_metadata, 'max_q_seqlen'): + updates['max_q_seqlen'] = 1 + if hasattr(token_attn_metadata, 'max_kv_seqlen'): + updates['max_kv_seqlen'] = getattr(token_attn_metadata, 'max_kv_seqlen') + if hasattr(token_attn_metadata, 'scheduler_metadata'): + updates['scheduler_metadata'] = None + if hasattr(token_attn_metadata, 'tile_scheduler_metadata'): + updates['tile_scheduler_metadata'] = None + if hasattr(token_attn_metadata, 'num_splits'): + updates['num_splits'] = None + if hasattr(token_attn_metadata, 'fill_seqlens'): + updates['fill_seqlens'] = concept_q_seqlens + return replace(token_attn_metadata, **updates) + + def _get_max_concepts_per_request(self, token_attn_metadata: Any, concept_q_seqlens: Tensor) -> int: + """Return per-request concept attention bound without hidden context + access.""" + max_q_seqlen = getattr(token_attn_metadata, 'max_q_seqlen', None) + if max_q_seqlen is not None: + return self.concept_count_from_seq_len(int(max_q_seqlen), self.chunk_size) + + # Test/direct-call fallback. Serving should use the scheduler-provided + # Python max_q_seqlen above, as in DSV4 metadata construction. + return int(concept_q_seqlens.max().item()) if concept_q_seqlens.numel() > 0 else 0 + + @staticmethod + def _build_prefill_token_layout(token_attn_metadata: Any, position_ids: Tensor) -> _PrefillTokenLayout: + """Build packed token-stream positions from prefill attention + metadata.""" + q_seqlens = token_attn_metadata.q_seqlens + q_start_loc = getattr(token_attn_metadata, 'q_start_loc', None) + if q_start_loc is None: + q_start_loc = F.pad(torch.cumsum(q_seqlens, dim=0, dtype=torch.int32), (1, 0))[:-1] + + total_tokens = int(position_ids.numel()) + + q_seqlens_long = q_seqlens.to(dtype=torch.long, device=position_ids.device) + q_start_loc_long = q_start_loc.to(dtype=torch.long, device=position_ids.device) + cu_q_seqlens = getattr(token_attn_metadata, 'cu_seqlens_q', None) + if cu_q_seqlens is None: + cu_q_seqlens = F.pad(torch.cumsum(q_seqlens, dim=0, dtype=torch.int32), (1, 0)) + cu_q_seqlens_long = cu_q_seqlens.to(dtype=torch.long, device=position_ids.device) + + token_ids = torch.arange(total_tokens, dtype=torch.long, device=position_ids.device) + token_seq = torch.searchsorted(cu_q_seqlens_long[1:], token_ids, right=True) + token_pos = token_ids - cu_q_seqlens_long[token_seq] + + return _PrefillTokenLayout( + q_seqlens=q_seqlens, + q_start_loc=q_start_loc, + q_seqlens_long=q_seqlens_long, + q_start_loc_long=q_start_loc_long, + token_seq=token_seq, + token_pos=token_pos, + total_tokens=total_tokens, + ) + + def _build_prefill_concept_layout(self, token_attn_metadata: Any, position_ids: Tensor, + token_layout: _PrefillTokenLayout) -> _PrefillConceptLayout: + """Build compact chunk-token positions used by concept attention.""" + concept_q_seqlens_long = self.concept_counts_from_q_seqlens(token_layout.q_seqlens_long, self.chunk_size) + concept_q_seqlens = concept_q_seqlens_long.to(dtype=token_layout.q_seqlens.dtype, + device=token_layout.q_seqlens.device) + concept_q_start_loc = F.pad(torch.cumsum(concept_q_seqlens, dim=0, dtype=torch.int32), (1, 0))[:-1] + concept_q_start_loc_long = concept_q_start_loc.to(dtype=torch.long, device=position_ids.device) + concept_cu_seqlens_long = F.pad(torch.cumsum(concept_q_seqlens_long, dim=0), (1, 0)) + + # TODO: remove this eager compact-size scalar read by moving ConceptLM + # concept-stream allocation/compaction into the engine/backend contract, + # like DSV4's precomputed metadata and Qwen3.5's state-cache metadata. + num_concepts_total = int(concept_cu_seqlens_long[-1].item()) + concept_ids = torch.arange(num_concepts_total, dtype=torch.long, device=position_ids.device) + concept_seq = torch.searchsorted(concept_cu_seqlens_long[1:], concept_ids, right=True) + local_concept_ids = concept_ids - concept_cu_seqlens_long[concept_seq] + concept_token_start = token_layout.q_start_loc_long[concept_seq] + local_concept_ids * self.chunk_size + concept_position_ids = position_ids[concept_token_start] + max_concepts_per_request = self._get_max_concepts_per_request(token_attn_metadata, concept_q_seqlens_long) + + return _PrefillConceptLayout( + q_seqlens=concept_q_seqlens, + q_seqlens_long=concept_q_seqlens_long, + q_start_loc=concept_q_start_loc, + q_start_loc_long=concept_q_start_loc_long, + seq=concept_seq, + local_ids=local_concept_ids, + position_ids=concept_position_ids, + num_total=num_concepts_total, + max_per_request=max_concepts_per_request, + ) + + def _build_prefill_repeat_ids(self, token_layout: _PrefillTokenLayout, + concept_layout: _PrefillConceptLayout) -> Tensor: + """Map packed token rows to the concept row visible after shift.""" + seq_concept_start = concept_layout.q_start_loc_long[token_layout.token_seq] + seq_concept_count = concept_layout.q_seqlens_long[token_layout.token_seq] + repeat_slots = self.repeat_slot_ids(token_layout.token_pos, self.chunk_size, self.shift_feature) + valid_repeat = (repeat_slots >= 0) & (repeat_slots < seq_concept_count) + return torch.where( + valid_repeat, + seq_concept_start + repeat_slots, + torch.full_like(repeat_slots, -1), + ) + + def _build_prefill_merge_layout(self, token_layout: _PrefillTokenLayout, + concept_layout: _PrefillConceptLayout) -> _PrefillMergeLayout: + """Build token-to-concept merge metadata for mean/first/last modes.""" + seq_concept_start = concept_layout.q_start_loc_long[token_layout.token_seq] + seq_concept_count = concept_layout.q_seqlens_long[token_layout.token_seq] + merge_slots = torch.div(token_layout.token_pos, self.chunk_size, rounding_mode='floor') + valid_merge = (merge_slots >= 0) & (merge_slots < seq_concept_count) + merge_token_to_concept = torch.where( + valid_merge, + seq_concept_start + merge_slots, + torch.full_like(merge_slots, -1), + ) + safe_merge_ids = merge_token_to_concept.clamp(min=0) + merge_token_counts = torch.zeros(concept_layout.num_total, + dtype=torch.int32, + device=token_layout.token_pos.device) + merge_token_counts.index_add_(0, safe_merge_ids, valid_merge.to(dtype=torch.int32)) + + concept_seq_len = token_layout.q_seqlens_long[concept_layout.seq] + merge_first_pos = torch.where( + concept_seq_len < self.chunk_size, + torch.zeros_like(concept_layout.local_ids), + concept_layout.local_ids * self.chunk_size, + ) + merge_last_pos = torch.where( + concept_seq_len < self.chunk_size, + (concept_seq_len - 1).clamp(min=0), + concept_layout.local_ids * self.chunk_size + self.chunk_size - 1, + ) + merge_first_token_ids = token_layout.q_start_loc_long[concept_layout.seq] + merge_first_pos + merge_last_token_ids = token_layout.q_start_loc_long[concept_layout.seq] + merge_last_pos + merge_short_concept_mask = concept_seq_len < self.chunk_size + + return _PrefillMergeLayout( + token_to_concept=merge_token_to_concept, + token_counts=merge_token_counts, + first_token_ids=merge_first_token_ids, + last_token_ids=merge_last_token_ids, + short_concept_mask=merge_short_concept_mask, + ) + + def build_prefill_metadata(self, token_attn_metadata: Any, position_ids: Tensor) -> ConceptPrefillMetadata: + """Build packed token-to-concept metadata for batched prefill.""" + token_layout = self._build_prefill_token_layout(token_attn_metadata, position_ids) + concept_layout = self._build_prefill_concept_layout(token_attn_metadata, position_ids, token_layout) + token_to_concept = self._build_prefill_repeat_ids(token_layout, concept_layout) + merge_layout = self._build_prefill_merge_layout(token_layout, concept_layout) + + return ConceptPrefillMetadata( + token_q_seqlens=token_layout.q_seqlens, + token_q_start_loc=token_layout.q_start_loc, + concept_q_seqlens=concept_layout.q_seqlens, + concept_q_start_loc=concept_layout.q_start_loc, + concept_position_ids=concept_layout.position_ids, + merge_token_to_concept=merge_layout.token_to_concept, + merge_token_counts=merge_layout.token_counts, + merge_first_token_ids=merge_layout.first_token_ids, + merge_last_token_ids=merge_layout.last_token_ids, + merge_short_concept_mask=merge_layout.short_concept_mask, + token_to_concept=token_to_concept, + num_tokens_total=token_layout.total_tokens, + num_concepts_total=concept_layout.num_total, + max_concepts_per_request=concept_layout.max_per_request, + ) + + @staticmethod + def _merge_chunks_mean_packed(hidden_states: Tensor, prefill_metadata: ConceptPrefillMetadata) -> Tensor: + """Mean-pool packed token states by precomputed concept ids.""" + merge_token_to_concept = prefill_metadata.merge_token_to_concept.to(device=hidden_states.device) + valid_merge = (merge_token_to_concept >= 0).to(dtype=hidden_states.dtype).unsqueeze(-1) + safe_merge_ids = merge_token_to_concept.clamp(min=0) + merged = hidden_states.new_zeros((prefill_metadata.num_concepts_total, hidden_states.size(-1))) + merged.index_add_(0, safe_merge_ids, hidden_states * valid_merge) + counts = prefill_metadata.merge_token_counts.clamp(min=1).to(device=hidden_states.device, + dtype=hidden_states.dtype) + return merged / counts.unsqueeze(-1) + + def merge_chunks_packed(self, hidden_states: Tensor, prefill_metadata: ConceptPrefillMetadata) -> Tensor: + """Merge packed token states into packed per-request concept states.""" + if self.merge_method == 'first': + merged = hidden_states[prefill_metadata.merge_first_token_ids] + short_mean = self._merge_chunks_mean_packed(hidden_states, prefill_metadata) + return torch.where(prefill_metadata.merge_short_concept_mask[:, None], short_mean, merged) + if self.merge_method == 'last': + merged = hidden_states[prefill_metadata.merge_last_token_ids] + short_mean = self._merge_chunks_mean_packed(hidden_states, prefill_metadata) + return torch.where(prefill_metadata.merge_short_concept_mask[:, None], short_mean, merged) + + return self._merge_chunks_mean_packed(hidden_states, prefill_metadata) + + @staticmethod + def _gather_zero_prefixed_concepts(concept_states_with_zero: Tensor, + prefill_metadata: ConceptPrefillMetadata) -> Tensor: + """Gather zero-prefixed concept rows to packed token rows.""" + token_to_concept = prefill_metadata.token_to_concept.to(device=concept_states_with_zero.device) + gather_ids = torch.clamp(token_to_concept + 1, min=0) + return concept_states_with_zero[gather_ids] + + def repeat_shift_packed(self, concept_states: Tensor, prefill_metadata: ConceptPrefillMetadata) -> Tensor: + """Gather packed concept states back to packed token states.""" + concept_states_with_zero = torch.cat((torch.zeros_like(concept_states[:1]), concept_states), dim=0) + return self._gather_zero_prefixed_concepts(concept_states_with_zero, prefill_metadata) + + def repeat_shift_source_states_packed(self, concept_states_with_zero: Tensor, + prefill_metadata: ConceptPrefillMetadata) -> Tensor: + """Gather zero-prefixed packed concept source states to token rows.""" + return self._gather_zero_prefixed_concepts(concept_states_with_zero, prefill_metadata) + + def build_concept_prefill_metadata(self, token_attn_metadata: Any, prefill_metadata: ConceptPrefillMetadata): + """Build chunk-stream attention metadata for packed prefill.""" + concept_q_seqlens = prefill_metadata.concept_q_seqlens + concept_q_start_loc = prefill_metadata.concept_q_start_loc + concept_cu_seqlens = F.pad( + torch.cumsum(concept_q_seqlens, dim=0, dtype=torch.int32), + (1, 0), + ) + max_concept_seqlen = int(prefill_metadata.max_concepts_per_request) + + updates = dict( + is_decoding=False, + q_start_loc=concept_q_start_loc, + q_seqlens=concept_q_seqlens, + kv_seqlens=concept_q_seqlens, + cu_seqlens_q=concept_cu_seqlens, + cu_seqlens_k=concept_cu_seqlens, + ) + if hasattr(token_attn_metadata, 'kv_start_loc'): + updates['kv_start_loc'] = concept_q_start_loc + if hasattr(token_attn_metadata, 'kv_flatten_size'): + updates['kv_flatten_size'] = int(prefill_metadata.num_concepts_total) + if hasattr(token_attn_metadata, 'max_q_seqlen'): + updates['max_q_seqlen'] = max_concept_seqlen + if hasattr(token_attn_metadata, 'max_kv_seqlen'): + updates['max_kv_seqlen'] = max_concept_seqlen + return replace(token_attn_metadata, **updates) + + def merge_prefill_tail_chunk_states(self, source_states: Tensor, + prefill_metadata: ConceptPrefillMetadata) -> Tensor: + """Build per-request partial chunk accumulator rows after prefill.""" + device = source_states.device + q_seqlens = prefill_metadata.token_q_seqlens.to(device=device, dtype=torch.long) + q_start_loc = prefill_metadata.token_q_start_loc.to(device=device, dtype=torch.long) + batch_size = q_seqlens.size(0) + tail_lens = torch.remainder(q_seqlens, self.chunk_size) + tail_lens = torch.where(q_seqlens < self.chunk_size, q_seqlens, tail_lens) + tail_lens = torch.where(q_seqlens > 0, tail_lens, torch.zeros_like(tail_lens)) + has_tail = tail_lens > 0 + + tail_rows = source_states.new_zeros((batch_size, source_states.size(1), source_states.size(2))) + if source_states.size(0) == 0: + return tail_rows + if self.merge_method == 'first': + token_ids = q_start_loc + q_seqlens - tail_lens + token_ids = token_ids.clamp(min=0, max=max(source_states.size(0) - 1, 0)) + rows = source_states[token_ids] + return torch.where(has_tail.view(batch_size, 1, 1), rows, tail_rows) + if self.merge_method == 'last': + token_ids = q_start_loc + q_seqlens - 1 + token_ids = token_ids.clamp(min=0, max=max(source_states.size(0) - 1, 0)) + rows = source_states[token_ids] + return torch.where(has_tail.view(batch_size, 1, 1), rows, tail_rows) + + token_ids = torch.arange(prefill_metadata.num_tokens_total, dtype=torch.long, device=device) + cu_q_seqlens = F.pad(torch.cumsum(q_seqlens, dim=0), (1, 0)) + token_seq = torch.searchsorted(cu_q_seqlens[1:], token_ids, right=True) + token_pos = token_ids - cu_q_seqlens[token_seq] + token_tail_start = q_seqlens[token_seq] - tail_lens[token_seq] + valid_tail = (tail_lens[token_seq] > 0) & (token_pos >= token_tail_start) + weighted_source = source_states * valid_tail.to(dtype=source_states.dtype).view(-1, 1, 1) + tail_rows.index_add_(0, token_seq, weighted_source) + return tail_rows + + @staticmethod + def stack_concept_raw_states(concept_raw_states: list[Tensor]) -> Tensor: + """Stack raw concept-layer states to ``[rows, concept_layers, + hidden]``.""" + return torch.stack(tuple(concept_raw_states), dim=1) + + def write_prefill_state_caches_eager( + self, + chunk_source_state: Tensor | None, + last_raw_states: Tensor | None, + last_final_state: Tensor | None, + state_ids: Tensor | None, + prefill_metadata: ConceptPrefillMetadata, + source_states: Tensor, + predicted_vectors: Tensor, + concept_raw_states: list[Tensor], + ) -> None: + """Seed decode state caches from a completed prefill forward.""" + if state_ids is None: + return + if chunk_source_state is None or last_raw_states is None or last_final_state is None: + return + + state_ids = state_ids.to(device=source_states.device, dtype=torch.long).reshape(-1) + valid_indices = torch.nonzero(state_ids >= 0, as_tuple=False).flatten() + if valid_indices.numel() == 0: + return + + valid_state_ids = state_ids.index_select(0, valid_indices) + tail_rows = self.merge_prefill_tail_chunk_states(source_states, prefill_metadata) + chunk_source_state.index_copy_(0, valid_state_ids, + tail_rows.index_select(0, valid_indices).to(dtype=chunk_source_state.dtype)) + + concept_counts = prefill_metadata.concept_q_seqlens.to(device=source_states.device, dtype=torch.long) + concept_start = prefill_metadata.concept_q_start_loc.to(device=source_states.device, dtype=torch.long) + concept_valid_mask = (state_ids >= 0) & (concept_counts > 0) + concept_indices = torch.nonzero(concept_valid_mask, as_tuple=False).flatten() + if concept_indices.numel() == 0: + return + + concept_state_ids = state_ids.index_select(0, concept_indices) + last_concept_ids = concept_start.index_select(0, concept_indices) + concept_counts.index_select( + 0, concept_indices) - 1 + last_final_rows = predicted_vectors.index_select(0, last_concept_ids) + last_final_state.index_copy_(0, concept_state_ids, last_final_rows.to(dtype=last_final_state.dtype)) + raw_rows = self.stack_concept_raw_states(concept_raw_states).index_select(0, last_concept_ids) + last_raw_states.index_copy_(0, concept_state_ids, raw_rows.to(dtype=last_raw_states.dtype)) + + def snapshot_decode_concept_kv(self, concept_past_key_values: list[list[Tensor]], + concept_attn_metadata: Any) -> list[tuple[Tensor, Tensor]]: + """Snapshot concept KV slots that dummy non-boundary rows may + overwrite.""" + return [ + self.decode_kv_cache_snapshot( + k_cache, + v_cache, + concept_attn_metadata.block_offsets, + concept_attn_metadata.kv_seqlens, + ) + for k_cache, v_cache in concept_past_key_values + ] + + def restore_decode_concept_kv(self, concept_past_key_values: list[list[Tensor]], concept_attn_metadata: Any, + saved_kv: list[tuple[Tensor, Tensor]], restore_mask: Tensor) -> None: + """Restore concept KV slots for non-boundary and padded rows.""" + for (k_cache, v_cache), (saved_k, saved_v) in zip(concept_past_key_values, saved_kv): + self.decode_kv_cache_restore( + k_cache, + v_cache, + saved_k, + saved_v, + concept_attn_metadata.block_offsets, + concept_attn_metadata.kv_seqlens, + restore_mask, + ) + + def write_decode_concept_states( + self, + last_raw_state_cache: Tensor, + last_final_state_cache: Tensor, + predicted_vectors: Tensor, + raw_states: list[Tensor], + state_ids: Tensor, + update_mask: Tensor, + ) -> None: + """Write newly emitted concept states to persistent decode caches.""" + raw_rows = self.stack_concept_raw_states(raw_states) + self.decode_concept_state_update( + last_raw_state_cache, + last_final_state_cache, + predicted_vectors, + raw_rows, + state_ids, + update_mask, + ) + @abstractmethod def decode_chunk_state_update( self, @@ -68,6 +661,6 @@ class ConceptLMRuntimeOpsBuilder(ABC): @staticmethod @abstractmethod - def build() -> ConceptLMRuntimeOpsImpl: + def build(config: PretrainedConfig) -> ConceptLMRuntimeOpsImpl: """Build layer implementation.""" raise NotImplementedError('Not implemented.') diff --git a/lmdeploy/pytorch/backends/cuda/conceptlm.py b/lmdeploy/pytorch/backends/cuda/conceptlm.py index 6957ba910e..5f9c55e29c 100644 --- a/lmdeploy/pytorch/backends/cuda/conceptlm.py +++ b/lmdeploy/pytorch/backends/cuda/conceptlm.py @@ -80,6 +80,6 @@ class TritonConceptLMRuntimeOpsBuilder(ConceptLMRuntimeOpsBuilder): """Triton ConceptLM runtime operation builder.""" @staticmethod - def build() -> ConceptLMRuntimeOpsImpl: + def build(config) -> ConceptLMRuntimeOpsImpl: """Build layer implementation.""" - return TritonConceptLMRuntimeOpsImpl() + return TritonConceptLMRuntimeOpsImpl(config) diff --git a/lmdeploy/pytorch/backends/default/conceptlm.py b/lmdeploy/pytorch/backends/default/conceptlm.py index c72906d5d0..d1c3a459bf 100644 --- a/lmdeploy/pytorch/backends/default/conceptlm.py +++ b/lmdeploy/pytorch/backends/default/conceptlm.py @@ -5,20 +5,6 @@ from ..conceptlm import ConceptLMRuntimeOpsBuilder, ConceptLMRuntimeOpsImpl -def _flatten_decode_position_ids(position_ids: Tensor, batch_size: int, device: torch.device) -> Tensor: - """Normalize decode position ids to one absolute position per batch row.""" - if position_ids.dim() == 0: - position_ids = position_ids.view(1) - if position_ids.dim() == 1: - return position_ids.to(device=device, dtype=torch.long) - position_ids = position_ids.reshape(-1) - if position_ids.numel() == batch_size: - return position_ids.to(device=device, dtype=torch.long) - assert position_ids.numel() % batch_size == 0, ( - f'Cannot map position_ids with {position_ids.numel()} elements to batch size {batch_size}.') - return position_ids.reshape(-1, batch_size)[-1].to(device=device, dtype=torch.long) - - class DefaultConceptLMRuntimeOpsImpl(ConceptLMRuntimeOpsImpl): """Torch fallback implementation of ConceptLM runtime operations.""" @@ -56,7 +42,7 @@ def decode_chunk_state_update( f'{tuple(chunk_source_state_cache.shape[1:])}.') state_ids = state_ids.to(device=current_source_states.device, dtype=torch.long) - position_ids = _flatten_decode_position_ids(position_ids, batch_size, current_source_states.device) + position_ids = self.flatten_decode_position_ids(position_ids, batch_size, current_source_states.device) assert position_ids.numel() == batch_size, ( f'Expected {batch_size} decode position ids, got {position_ids.numel()}.') @@ -144,6 +130,6 @@ class DefaultConceptLMRuntimeOpsBuilder(ConceptLMRuntimeOpsBuilder): """Torch fallback ConceptLM runtime operation builder.""" @staticmethod - def build() -> ConceptLMRuntimeOpsImpl: + def build(config) -> ConceptLMRuntimeOpsImpl: """Build layer implementation.""" - return DefaultConceptLMRuntimeOpsImpl() + return DefaultConceptLMRuntimeOpsImpl(config) diff --git a/lmdeploy/pytorch/models/intern_ncp/metadata.py b/lmdeploy/pytorch/models/intern_ncp/metadata.py index c7a55830a9..cd1dc7f045 100644 --- a/lmdeploy/pytorch/models/intern_ncp/metadata.py +++ b/lmdeploy/pytorch/models/intern_ncp/metadata.py @@ -173,79 +173,6 @@ def last_final_state(self) -> torch.Tensor | None: return self.semantic_state_cache(self.last_final_name, self.last_final_idx) -@dataclass -class ConceptChunkStateUpdateResult: - """Fixed-shape result of one decode chunk-source state update.""" - - concept_input_states: torch.Tensor - concept_update_mask: torch.Tensor - - -@dataclass -class ConceptDecodeMetadata: - """Fixed-layout decode metadata derived once from engine inputs. - - Decode is always represented as the engine's fixed ``[1, batch]`` token - layout at the model boundary and flattened to ``[batch]`` / ``[batch, H]`` - only inside ConceptLM helpers. - """ - - position_ids: torch.Tensor - state_ids: torch.Tensor - safe_state_ids: torch.Tensor - valid_state_mask: torch.Tensor - - -@dataclass -class ConceptPrefillMetadata: - """Packed prefill metadata derived once from token attention metadata. - - Field groups: - - token stream: original engine-provided token request boundaries. - - concept stream: compact chunk-token request boundaries and positions - used by the concept predictor attention. - - chunk merge: token -> concept ids and helper ids used to reduce encoder - token states into concept states without a Python batch loop. - - repeat/gather: concept -> token ids used to project compact concept - states back to the packed token stream. - - scalar bounds: eager compact sizes / upper bounds needed by metadata and - attention launch parameters. - """ - - # Token stream metadata, shape [batch]. This is the original packed prefill - # layout consumed by normal token attention. - token_q_seqlens: torch.Tensor - token_q_start_loc: torch.Tensor - - # Concept stream metadata, shape [batch] plus compact concept positions. - # These describe the shorter chunk-token stream consumed by concept - # predictor attention. - concept_q_seqlens: torch.Tensor - concept_q_start_loc: torch.Tensor - concept_position_ids: torch.Tensor - - # Chunk merge metadata. ``merge_token_to_concept`` maps each packed token to - # the compact concept row that owns it, or -1 when the token is dropped from - # concept production. Counts/first/last ids implement mean/first/last merge. - merge_token_to_concept: torch.Tensor - merge_token_counts: torch.Tensor - merge_first_token_ids: torch.Tensor - merge_last_token_ids: torch.Tensor - merge_short_concept_mask: torch.Tensor - - # Repeat/gather metadata. Maps each packed token row to the compact concept - # row it should read after shift semantics are applied, or -1 for the - # zero-concept row. - token_to_concept: torch.Tensor - - # Scalar sizes/bounds. ``num_concepts_total`` is the exact compact size in - # eager prefill; ``max_concepts_per_request`` is the per-request attention - # launch bound. - num_tokens_total: int - num_concepts_total: int - max_concepts_per_request: int - - def _split_concept_past_key_values(config: PretrainedConfig, past_key_values: list[list[torch.Tensor]] | None): """Split the flat LMDeploy KV-cache list into ConceptLM streams.""" if past_key_values is None or len(past_key_values) == 0: @@ -262,17 +189,3 @@ def _split_concept_past_key_values(config: PretrainedConfig, past_key_values: li concept_end = enc_end + concept_layers dec_end = concept_end + dec_layers return past_key_values[:enc_end], past_key_values[enc_end:concept_end], past_key_values[concept_end:dec_end] - - -def _flatten_decode_position_ids(position_ids: torch.Tensor, batch_size: int) -> torch.Tensor: - """Normalize decode position ids to one absolute position per batch row.""" - if position_ids.dim() == 0: - position_ids = position_ids.view(1) - if position_ids.dim() == 1: - return position_ids.to(torch.long) - position_ids = position_ids.reshape(-1) - if position_ids.numel() == batch_size: - return position_ids.to(torch.long) - assert position_ids.numel() % batch_size == 0, ( - f'Cannot map position_ids with {position_ids.numel()} elements to batch size {batch_size}.') - return position_ids.reshape(-1, batch_size)[-1].to(torch.long) diff --git a/lmdeploy/pytorch/models/intern_ncp/modeling.py b/lmdeploy/pytorch/models/intern_ncp/modeling.py index ce6393c498..702aec35be 100644 --- a/lmdeploy/pytorch/models/intern_ncp/modeling.py +++ b/lmdeploy/pytorch/models/intern_ncp/modeling.py @@ -5,9 +5,13 @@ import torch from torch import nn -from torch.nn import functional as F from transformers.configuration_utils import PretrainedConfig +from lmdeploy.pytorch.backends.conceptlm import ( + ConceptChunkStateUpdateResult, + ConceptDecodeMetadata, + ConceptPrefillMetadata, +) from lmdeploy.pytorch.model_inputs import StepContext, StepContextManager from lmdeploy.pytorch.nn import ConceptLMRuntimeOps from lmdeploy.pytorch.weight_loader.model_weight_loader import load_weight @@ -17,11 +21,7 @@ from ..utils.model import DeployModelMixinV1 from .metadata import ( ConceptCaches, - ConceptChunkStateUpdateResult, - ConceptDecodeMetadata, ConceptMetadata, - ConceptPrefillMetadata, - _flatten_decode_position_ids, ) from .modules import ( ConceptPredictor, @@ -93,45 +93,6 @@ class _EncoderOutput: history_buffer: torch.Tensor -@dataclass -class _PrefillTokenLayout: - """Packed token-stream layout derived from prefill attention metadata.""" - - q_seqlens: torch.Tensor - q_start_loc: torch.Tensor - q_seqlens_long: torch.Tensor - q_start_loc_long: torch.Tensor - token_seq: torch.Tensor - token_pos: torch.Tensor - total_tokens: int - - -@dataclass -class _PrefillConceptLayout: - """Compact chunk-token stream layout used by concept predictor prefill.""" - - q_seqlens: torch.Tensor - q_seqlens_long: torch.Tensor - q_start_loc: torch.Tensor - q_start_loc_long: torch.Tensor - seq: torch.Tensor - local_ids: torch.Tensor - position_ids: torch.Tensor - num_total: int - max_per_request: int - - -@dataclass -class _PrefillMergeLayout: - """Token-to-concept merge metadata for compact prefill.""" - - token_to_concept: torch.Tensor - token_counts: torch.Tensor - first_token_ids: torch.Tensor - last_token_ids: torch.Tensor - short_concept_mask: torch.Tensor - - class ConceptLMV22VQForCausalLM(nn.Module, DeployModelMixinV1, CudaGraphMixin): """Rewrote model of ConceptLMV22VQForCausalLM.""" @@ -211,7 +172,7 @@ def __init__(self, torch.zeros(config.concept_decoder_layers, 2, dtype=dtype, device=device), requires_grad=False) self.dd_two_route_add = TwoRouteAdd(config, dtype=dtype, device=device) - self.concept_ops = ConceptLMRuntimeOps() + self.concept_ops = ConceptLMRuntimeOps(config) # output projection — mirrors ``output_layer`` in the reference. Built # via build_lm_head and named ``lm_head`` so DeployModelMixinV1. # get_logits picks it up directly; load_weights maps the checkpoint's @@ -361,64 +322,20 @@ def _normalize_decode_inputs(self, hidden_states: torch.Tensor, position_ids: to f'ConceptLM decode expects fixed engine layout [1, batch, hidden], ' f'got {tuple(hidden_states.shape)}.') batch_size = hidden_states.size(1) - position_ids = _flatten_decode_position_ids(position_ids, batch_size).to(device=hidden_states.device) + position_ids = self.concept_ops.flatten_decode_position_ids(position_ids, batch_size, hidden_states.device) if position_ids.numel() != batch_size: raise ValueError(f'Expected {batch_size} decode position ids, got {position_ids.numel()}.') return hidden_states[0].contiguous(), position_ids - @staticmethod - def _build_decode_metadata(position_ids: torch.Tensor, - state_ids: torch.Tensor | None, - batch_size: int, - device: torch.device) -> ConceptDecodeMetadata: - """Build fixed-shape decode metadata from engine state ids.""" - if state_ids is None: - raise RuntimeError('ConceptLM decode requires state_ids.') - state_ids = state_ids.to(device=device, dtype=torch.long).reshape(-1) - if state_ids.numel() != batch_size: - raise ValueError(f'Expected {batch_size} decode state ids, got {state_ids.numel()}.') - valid_state_mask = state_ids >= 0 - return ConceptDecodeMetadata( - position_ids=position_ids, - state_ids=state_ids, - safe_state_ids=state_ids.clamp(min=0), - valid_state_mask=valid_state_mask, - ) - - @staticmethod - def _select_decode_state_rows(state_cache: torch.Tensor, - decode_metadata: ConceptDecodeMetadata) -> torch.Tensor: - """Gather state-cache rows and zero out padded decode rows.""" - rows = state_cache.index_select(0, decode_metadata.safe_state_ids) - mask_shape = (decode_metadata.valid_state_mask.size(0), ) + (1, ) * (rows.dim() - 1) - valid_mask = decode_metadata.valid_state_mask.view(mask_shape) - return torch.where(valid_mask, rows, torch.zeros_like(rows)) - def _select_decode_last_state_rows(self, concept_caches: ConceptCaches, decode_metadata: ConceptDecodeMetadata) -> tuple[torch.Tensor, torch.Tensor]: - """Gather packed last-concept state rows once and return final/raw - views.""" - last_state = concept_caches.last_state - if last_state is not None: - rows = self._select_decode_state_rows(last_state, decode_metadata) - return rows[:, 0], rows[:, 1:] - - last_final_state = concept_caches.last_final_state - last_raw_states = concept_caches.last_raw_states - return ( - self._select_decode_state_rows(last_final_state, decode_metadata), - self._select_decode_state_rows(last_raw_states, decode_metadata), - ) - - def _decode_concept_read_mask(self, decode_metadata: ConceptDecodeMetadata) -> torch.Tensor: - """Return rows whose current decode token should read a cached - concept.""" - repeat_slots = self._repeat_slot_ids( - decode_metadata.position_ids, - int(self.config.concept_chunk_size), - bool(getattr(self.config, 'concept_shift_feature', True)), + """Gather latest concept state rows through backend-owned layout.""" + return self.concept_ops.select_decode_last_state_rows( + concept_caches.last_state, + concept_caches.last_final_state, + concept_caches.last_raw_states, + decode_metadata, ) - return decode_metadata.valid_state_mask & (repeat_slots >= 0) def _select_decode_decoder_concepts(self, concept_caches: ConceptCaches, @@ -432,7 +349,7 @@ def _select_decode_decoder_concepts(self, final_state = previous_final_state raw_states = previous_raw_states - concept_read_mask = self._decode_concept_read_mask(decode_metadata) + concept_read_mask = self.concept_ops.decode_concept_read_mask(decode_metadata) final_state = torch.where(concept_read_mask.view(-1, 1), final_state, torch.zeros_like(final_state)) raw_states = torch.where(concept_read_mask.view(-1, 1, 1), raw_states, torch.zeros_like(raw_states)) return _DecoderConceptInput.for_decode(final_state, raw_states) @@ -455,83 +372,15 @@ def _decode_concept_position_ids(position_ids: torch.Tensor, chunk_size: int) -> boundaries.""" return (position_ids - int(chunk_size) + 1).clamp(min=0) - def _build_concept_decode_metadata_static(self, - token_attn_metadata: Any, - decode_metadata: ConceptDecodeMetadata): - """Build fixed-shape concept-stream decode metadata. - - Concept predictor runs with the same batch shape as token decode. - Boundary rows append a real concept KV entry. Non-boundary/padded rows - execute dummy concept attention with safe ``kv_seqlens >= 1``; their KV - writes are restored afterward by ``ConceptLMRuntimeOps``. - """ - device = decode_metadata.position_ids.device - batch_size = decode_metadata.position_ids.numel() - q_seqlens = token_attn_metadata.q_seqlens - q_start_loc = token_attn_metadata.q_start_loc - kv_seqlens = token_attn_metadata.kv_seqlens - q_dtype = q_seqlens.dtype - q_start_dtype = q_start_loc.dtype - kv_dtype = kv_seqlens.dtype - - concept_q_seqlens = torch.ones((batch_size, ), dtype=q_dtype, device=device) - concept_q_start_loc = torch.arange(batch_size, dtype=q_start_dtype, device=device) - concept_cu_seqlens = F.pad(torch.cumsum(concept_q_seqlens, dim=0, dtype=torch.int32), (1, 0)) - concept_kv_seqlens = torch.div( - decode_metadata.position_ids + 1, - int(self.config.concept_chunk_size), - rounding_mode='floor', - ).clamp(min=1).to(dtype=kv_dtype) - - updates = dict( - is_decoding=True, - block_offsets=token_attn_metadata.block_offsets, - q_start_loc=concept_q_start_loc, - q_seqlens=concept_q_seqlens, - kv_seqlens=concept_kv_seqlens, - cu_seqlens_q=concept_cu_seqlens, - cu_seqlens_k=concept_cu_seqlens, - ) - if hasattr(token_attn_metadata, 'kv_start_loc'): - updates['kv_start_loc'] = concept_kv_seqlens - concept_q_seqlens.to(dtype=concept_kv_seqlens.dtype) - if hasattr(token_attn_metadata, 'kv_flatten_size'): - updates['kv_flatten_size'] = batch_size - if hasattr(token_attn_metadata, 'max_q_seqlen'): - updates['max_q_seqlen'] = 1 - if hasattr(token_attn_metadata, 'max_kv_seqlen'): - # Decode kernels use kv_seqlens directly. Keep a conservative bound - # without reading the dynamic maximum back to host. - updates['max_kv_seqlen'] = getattr(token_attn_metadata, 'max_kv_seqlen') - if hasattr(token_attn_metadata, 'scheduler_metadata'): - updates['scheduler_metadata'] = None - if hasattr(token_attn_metadata, 'tile_scheduler_metadata'): - updates['tile_scheduler_metadata'] = None - if hasattr(token_attn_metadata, 'num_splits'): - updates['num_splits'] = None - if hasattr(token_attn_metadata, 'fill_seqlens'): - updates['fill_seqlens'] = concept_q_seqlens - return replace(token_attn_metadata, **updates) - - @staticmethod - def _stack_concept_raw_states(concept_raw_states: list[torch.Tensor]) -> torch.Tensor: - """Stack raw concept-layer states to ``[rows, concept_layers, - hidden]``.""" - return torch.stack(tuple(concept_raw_states), dim=1) - def _snapshot_decode_concept_kv(self, concept_caches: ConceptCaches, concept_attn_metadata: Any): """Snapshot concept KV slots that dummy non-boundary rows may overwrite.""" - return [ - self.concept_ops.decode_kv_cache_snapshot( - k_cache, - v_cache, - concept_attn_metadata.block_offsets, - concept_attn_metadata.kv_seqlens, - ) - for k_cache, v_cache in concept_caches.concept_past_key_values - ] + return self.concept_ops.snapshot_decode_concept_kv( + concept_caches.concept_past_key_values, + concept_attn_metadata, + ) def _restore_decode_concept_kv_( self, @@ -541,16 +390,12 @@ def _restore_decode_concept_kv_( restore_mask: torch.Tensor, ): """Restore concept KV slots for non-boundary and padded rows.""" - for (k_cache, v_cache), (saved_k, saved_v) in zip(concept_caches.concept_past_key_values, saved_kv): - self.concept_ops.decode_kv_cache_restore( - k_cache, - v_cache, - saved_k, - saved_v, - concept_attn_metadata.block_offsets, - concept_attn_metadata.kv_seqlens, - restore_mask, - ) + self.concept_ops.restore_decode_concept_kv( + concept_caches.concept_past_key_values, + concept_attn_metadata, + saved_kv, + restore_mask, + ) def _write_decode_concept_states_static_(self, concept_caches: ConceptCaches, @@ -559,14 +404,11 @@ def _write_decode_concept_states_static_(self, predicted_vectors: torch.Tensor, concept_raw_states: list[torch.Tensor]): """Write newly emitted concept states to persistent decode caches.""" - last_final_state = concept_caches.last_final_state - last_raw_states = concept_caches.last_raw_states - raw_rows = self._stack_concept_raw_states(concept_raw_states) - self.concept_ops.decode_concept_state_update( - last_raw_states, - last_final_state, + self.concept_ops.write_decode_concept_states( + concept_caches.last_raw_states, + concept_caches.last_final_state, predicted_vectors, - raw_rows, + concept_raw_states, decode_metadata.state_ids, update_mask, ) @@ -583,7 +425,7 @@ def _build_decode_concept_request(self, decode_metadata.position_ids, concept_metadata.chunk_size, ) - concept_attn_metadata = self._build_concept_decode_metadata_static( + concept_attn_metadata = self.concept_ops.build_concept_decode_metadata_static( concept_metadata.attn_metadata, decode_metadata, ) @@ -621,350 +463,12 @@ def _update_decode_concept_states_static_(self, concept_output.raw_states, ) - def _merge_prefill_tail_chunk_states(self, - source_states: torch.Tensor, - prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: - """Build per-request partial chunk accumulator rows after prefill.""" - device = source_states.device - chunk_size = int(self.config.concept_chunk_size) - q_seqlens = prefill_metadata.token_q_seqlens.to(device=device, dtype=torch.long) - q_start_loc = prefill_metadata.token_q_start_loc.to(device=device, dtype=torch.long) - batch_size = q_seqlens.size(0) - tail_lens = torch.remainder(q_seqlens, chunk_size) - tail_lens = torch.where(q_seqlens < chunk_size, q_seqlens, tail_lens) - tail_lens = torch.where(q_seqlens > 0, tail_lens, torch.zeros_like(tail_lens)) - has_tail = tail_lens > 0 - - tail_rows = source_states.new_zeros((batch_size, source_states.size(1), source_states.size(2))) - if source_states.size(0) == 0: - return tail_rows - merge_method = getattr(self.config, 'concept_chunk_merge_method', 'meanpooling') - if merge_method == 'first': - token_ids = q_start_loc + q_seqlens - tail_lens - token_ids = token_ids.clamp(min=0, max=max(source_states.size(0) - 1, 0)) - rows = source_states[token_ids] - return torch.where(has_tail.view(batch_size, 1, 1), rows, tail_rows) - if merge_method == 'last': - token_ids = q_start_loc + q_seqlens - 1 - token_ids = token_ids.clamp(min=0, max=max(source_states.size(0) - 1, 0)) - rows = source_states[token_ids] - return torch.where(has_tail.view(batch_size, 1, 1), rows, tail_rows) - - token_ids = torch.arange(prefill_metadata.num_tokens_total, dtype=torch.long, device=device) - cu_q_seqlens = F.pad(torch.cumsum(q_seqlens, dim=0), (1, 0)) - token_seq = torch.searchsorted(cu_q_seqlens[1:], token_ids, right=True) - token_pos = token_ids - cu_q_seqlens[token_seq] - token_tail_start = q_seqlens[token_seq] - tail_lens[token_seq] - valid_tail = (tail_lens[token_seq] > 0) & (token_pos >= token_tail_start) - weighted_source = source_states * valid_tail.to(dtype=source_states.dtype).view(-1, 1, 1) - tail_rows.index_add_(0, token_seq, weighted_source) - return tail_rows - - def _write_prefill_state_caches_eager_(self, - concept_caches: ConceptCaches, - concept_metadata: ConceptMetadata, - prefill_metadata: ConceptPrefillMetadata, - source_states: torch.Tensor, - predicted_vectors: torch.Tensor, - concept_raw_states: list[torch.Tensor]): - """Seed decode state caches from a completed prefill forward.""" - if concept_metadata.state_ids is None: - return - if concept_caches.state_caches is None and concept_caches.named_state_caches is None: - return - chunk_source_state = concept_caches.chunk_source_state - last_raw_states = concept_caches.last_raw_states - last_final_state = concept_caches.last_final_state - - state_ids = concept_metadata.state_ids.to(device=source_states.device, dtype=torch.long).reshape(-1) - valid_indices = torch.nonzero(state_ids >= 0, as_tuple=False).flatten() - if valid_indices.numel() == 0: - return - - valid_state_ids = state_ids.index_select(0, valid_indices) - tail_rows = self._merge_prefill_tail_chunk_states(source_states, prefill_metadata) - chunk_source_state.index_copy_(0, valid_state_ids, - tail_rows.index_select(0, valid_indices).to(dtype=chunk_source_state.dtype)) - - concept_counts = prefill_metadata.concept_q_seqlens.to(device=source_states.device, dtype=torch.long) - concept_start = prefill_metadata.concept_q_start_loc.to(device=source_states.device, dtype=torch.long) - concept_valid_mask = (state_ids >= 0) & (concept_counts > 0) - concept_indices = torch.nonzero(concept_valid_mask, as_tuple=False).flatten() - if concept_indices.numel() == 0: - return - - concept_state_ids = state_ids.index_select(0, concept_indices) - last_concept_ids = concept_start.index_select(0, concept_indices) + concept_counts.index_select( - 0, concept_indices) - 1 - last_final_rows = predicted_vectors.index_select(0, last_concept_ids) - last_final_state.index_copy_(0, concept_state_ids, last_final_rows.to(dtype=last_final_state.dtype)) - raw_rows = self._stack_concept_raw_states(concept_raw_states).index_select(0, last_concept_ids) - last_raw_states.index_copy_(0, concept_state_ids, raw_rows.to(dtype=last_raw_states.dtype)) - - @staticmethod - def _concept_count_from_seq_len(seq_len: int, chunk_size: int) -> int: - """Return reference ConceptLM chunk count for one request length.""" - seq_len = int(seq_len or 0) - if seq_len <= 0: - return 0 - if seq_len < chunk_size: - return 1 - return seq_len // chunk_size - - @staticmethod - def _concept_counts_from_q_seqlens(q_seqlens: torch.Tensor, chunk_size: int) -> torch.Tensor: - """Vectorized version of ``_concept_count_from_seq_len``.""" - counts = torch.div(q_seqlens, chunk_size, rounding_mode='floor').clamp(min=1) - return torch.where(q_seqlens > 0, counts, torch.zeros_like(q_seqlens)) - - @staticmethod - def _repeat_slot_ids(token_pos: torch.Tensor, chunk_size: int, shift_feature: bool) -> torch.Tensor: - """Return local concept slot read by each token after shift - semantics.""" - if shift_feature: - return torch.div(token_pos + 1, chunk_size, rounding_mode='floor') - 1 - return torch.div(token_pos, chunk_size, rounding_mode='floor') - 1 - - @staticmethod - def _get_max_concepts_per_request(token_attn_metadata: Any, - concept_q_seqlens: torch.Tensor, - chunk_size: int) -> int: - """Return per-request concept attention bound without hidden context - access.""" - max_q_seqlen = getattr(token_attn_metadata, 'max_q_seqlen', None) - if max_q_seqlen is not None: - return ConceptLMV22VQForCausalLM._concept_count_from_seq_len(int(max_q_seqlen), chunk_size) - - # Test/direct-call fallback. Serving should use the scheduler-provided - # Python max_q_seqlen above, as in DSV4 metadata construction. - return int(concept_q_seqlens.max().item()) if concept_q_seqlens.numel() > 0 else 0 - - @staticmethod - def _build_prefill_token_layout(token_attn_metadata: Any, - position_ids: torch.Tensor) -> _PrefillTokenLayout: - """Build packed token-stream positions from prefill attention - metadata.""" - q_seqlens = token_attn_metadata.q_seqlens - q_start_loc = getattr(token_attn_metadata, 'q_start_loc', None) - if q_start_loc is None: - q_start_loc = F.pad(torch.cumsum(q_seqlens, dim=0, dtype=torch.int32), (1, 0))[:-1] - - total_tokens = int(position_ids.numel()) - - q_seqlens_long = q_seqlens.to(dtype=torch.long, device=position_ids.device) - q_start_loc_long = q_start_loc.to(dtype=torch.long, device=position_ids.device) - cu_q_seqlens = getattr(token_attn_metadata, 'cu_seqlens_q', None) - if cu_q_seqlens is None: - cu_q_seqlens = F.pad(torch.cumsum(q_seqlens, dim=0, dtype=torch.int32), (1, 0)) - cu_q_seqlens_long = cu_q_seqlens.to(dtype=torch.long, device=position_ids.device) - - token_ids = torch.arange(total_tokens, dtype=torch.long, device=position_ids.device) - token_seq = torch.searchsorted(cu_q_seqlens_long[1:], token_ids, right=True) - token_pos = token_ids - cu_q_seqlens_long[token_seq] - - return _PrefillTokenLayout( - q_seqlens=q_seqlens, - q_start_loc=q_start_loc, - q_seqlens_long=q_seqlens_long, - q_start_loc_long=q_start_loc_long, - token_seq=token_seq, - token_pos=token_pos, - total_tokens=total_tokens, - ) - - def _build_prefill_concept_layout(self, - token_attn_metadata: Any, - position_ids: torch.Tensor, - token_layout: _PrefillTokenLayout, - chunk_size: int) -> _PrefillConceptLayout: - """Build compact chunk-token positions used by concept attention.""" - concept_q_seqlens_long = self._concept_counts_from_q_seqlens(token_layout.q_seqlens_long, chunk_size) - concept_q_seqlens = concept_q_seqlens_long.to(dtype=token_layout.q_seqlens.dtype, - device=token_layout.q_seqlens.device) - concept_q_start_loc = F.pad(torch.cumsum(concept_q_seqlens, dim=0, dtype=torch.int32), (1, 0))[:-1] - concept_q_start_loc_long = concept_q_start_loc.to(dtype=torch.long, device=position_ids.device) - concept_cu_seqlens_long = F.pad(torch.cumsum(concept_q_seqlens_long, dim=0), (1, 0)) - - # TODO: remove this eager compact-size scalar read by moving ConceptLM - # concept-stream allocation/compaction into the engine/backend contract, - # like DSV4's precomputed metadata and Qwen3.5's state-cache metadata. - num_concepts_total = int(concept_cu_seqlens_long[-1].item()) - concept_ids = torch.arange(num_concepts_total, dtype=torch.long, device=position_ids.device) - concept_seq = torch.searchsorted(concept_cu_seqlens_long[1:], concept_ids, right=True) - local_concept_ids = concept_ids - concept_cu_seqlens_long[concept_seq] - concept_token_start = token_layout.q_start_loc_long[concept_seq] + local_concept_ids * chunk_size - concept_position_ids = position_ids[concept_token_start] - max_concepts_per_request = self._get_max_concepts_per_request( - token_attn_metadata, - concept_q_seqlens_long, - chunk_size, - ) - - return _PrefillConceptLayout( - q_seqlens=concept_q_seqlens, - q_seqlens_long=concept_q_seqlens_long, - q_start_loc=concept_q_start_loc, - q_start_loc_long=concept_q_start_loc_long, - seq=concept_seq, - local_ids=local_concept_ids, - position_ids=concept_position_ids, - num_total=num_concepts_total, - max_per_request=max_concepts_per_request, - ) - - def _build_prefill_repeat_ids(self, - token_layout: _PrefillTokenLayout, - concept_layout: _PrefillConceptLayout, - chunk_size: int, - shift_feature: bool) -> torch.Tensor: - """Map packed token rows to the concept row visible after shift.""" - seq_concept_start = concept_layout.q_start_loc_long[token_layout.token_seq] - seq_concept_count = concept_layout.q_seqlens_long[token_layout.token_seq] - repeat_slots = self._repeat_slot_ids(token_layout.token_pos, chunk_size, shift_feature) - valid_repeat = (repeat_slots >= 0) & (repeat_slots < seq_concept_count) - return torch.where( - valid_repeat, - seq_concept_start + repeat_slots, - torch.full_like(repeat_slots, -1), - ) - - @staticmethod - def _build_prefill_merge_layout(token_layout: _PrefillTokenLayout, - concept_layout: _PrefillConceptLayout, - chunk_size: int) -> _PrefillMergeLayout: - """Build token-to-concept merge metadata for mean/first/last modes.""" - seq_concept_start = concept_layout.q_start_loc_long[token_layout.token_seq] - seq_concept_count = concept_layout.q_seqlens_long[token_layout.token_seq] - merge_slots = torch.div(token_layout.token_pos, chunk_size, rounding_mode='floor') - valid_merge = (merge_slots >= 0) & (merge_slots < seq_concept_count) - merge_token_to_concept = torch.where( - valid_merge, - seq_concept_start + merge_slots, - torch.full_like(merge_slots, -1), - ) - safe_merge_ids = merge_token_to_concept.clamp(min=0) - merge_token_counts = torch.zeros(concept_layout.num_total, - dtype=torch.int32, - device=token_layout.token_pos.device) - merge_token_counts.index_add_(0, safe_merge_ids, valid_merge.to(dtype=torch.int32)) - - concept_seq_len = token_layout.q_seqlens_long[concept_layout.seq] - merge_first_pos = torch.where( - concept_seq_len < chunk_size, - torch.zeros_like(concept_layout.local_ids), - concept_layout.local_ids * chunk_size, - ) - merge_last_pos = torch.where( - concept_seq_len < chunk_size, - (concept_seq_len - 1).clamp(min=0), - concept_layout.local_ids * chunk_size + chunk_size - 1, - ) - merge_first_token_ids = token_layout.q_start_loc_long[concept_layout.seq] + merge_first_pos - merge_last_token_ids = token_layout.q_start_loc_long[concept_layout.seq] + merge_last_pos - merge_short_concept_mask = concept_seq_len < chunk_size - - return _PrefillMergeLayout( - token_to_concept=merge_token_to_concept, - token_counts=merge_token_counts, - first_token_ids=merge_first_token_ids, - last_token_ids=merge_last_token_ids, - short_concept_mask=merge_short_concept_mask, - ) - - def _build_prefill_metadata(self, - token_attn_metadata: Any, - position_ids: torch.Tensor) -> ConceptPrefillMetadata: - """Build packed token-to-concept metadata for batched prefill.""" - chunk_size = int(self.config.concept_chunk_size) - shift_feature = bool(getattr(self.config, 'concept_shift_feature', True)) - token_layout = self._build_prefill_token_layout(token_attn_metadata, position_ids) - concept_layout = self._build_prefill_concept_layout( - token_attn_metadata, - position_ids, - token_layout, - chunk_size, - ) - token_to_concept = self._build_prefill_repeat_ids( - token_layout, - concept_layout, - chunk_size, - shift_feature, - ) - merge_layout = self._build_prefill_merge_layout(token_layout, concept_layout, chunk_size) - - return ConceptPrefillMetadata( - token_q_seqlens=token_layout.q_seqlens, - token_q_start_loc=token_layout.q_start_loc, - concept_q_seqlens=concept_layout.q_seqlens, - concept_q_start_loc=concept_layout.q_start_loc, - concept_position_ids=concept_layout.position_ids, - merge_token_to_concept=merge_layout.token_to_concept, - merge_token_counts=merge_layout.token_counts, - merge_first_token_ids=merge_layout.first_token_ids, - merge_last_token_ids=merge_layout.last_token_ids, - merge_short_concept_mask=merge_layout.short_concept_mask, - token_to_concept=token_to_concept, - num_tokens_total=token_layout.total_tokens, - num_concepts_total=concept_layout.num_total, - max_concepts_per_request=concept_layout.max_per_request, - ) - - @staticmethod - def _merge_chunks_mean_packed(hidden_states: torch.Tensor, - prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: - """Mean-pool packed token states by precomputed concept ids.""" - merge_token_to_concept = prefill_metadata.merge_token_to_concept.to(device=hidden_states.device) - valid_merge = (merge_token_to_concept >= 0).to(dtype=hidden_states.dtype).unsqueeze(-1) - safe_merge_ids = merge_token_to_concept.clamp(min=0) - merged = hidden_states.new_zeros((prefill_metadata.num_concepts_total, hidden_states.size(-1))) - merged.index_add_(0, safe_merge_ids, hidden_states * valid_merge) - counts = prefill_metadata.merge_token_counts.clamp(min=1).to(device=hidden_states.device, - dtype=hidden_states.dtype) - return merged / counts.unsqueeze(-1) - - def _merge_chunks_packed(self, - hidden_states: torch.Tensor, - prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: - """Merge packed token states into packed per-request concept states.""" - merge_method = getattr(self.config, 'concept_chunk_merge_method', 'meanpooling') - if merge_method == 'first': - merged = hidden_states[prefill_metadata.merge_first_token_ids] - short_mean = self._merge_chunks_mean_packed(hidden_states, prefill_metadata) - return torch.where(prefill_metadata.merge_short_concept_mask[:, None], short_mean, merged) - if merge_method == 'last': - merged = hidden_states[prefill_metadata.merge_last_token_ids] - short_mean = self._merge_chunks_mean_packed(hidden_states, prefill_metadata) - return torch.where(prefill_metadata.merge_short_concept_mask[:, None], short_mean, merged) - - return self._merge_chunks_mean_packed(hidden_states, prefill_metadata) - - @staticmethod - def _gather_zero_prefixed_concepts(concept_states_with_zero: torch.Tensor, - prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: - """Gather zero-prefixed concept rows to packed token rows.""" - token_to_concept = prefill_metadata.token_to_concept.to(device=concept_states_with_zero.device) - gather_ids = torch.clamp(token_to_concept + 1, min=0) - return concept_states_with_zero[gather_ids] - - def _repeat_shift_packed(self, - concept_states: torch.Tensor, - prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: - """Gather packed concept states back to packed token states.""" - concept_states_with_zero = torch.cat((torch.zeros_like(concept_states[:1]), concept_states), dim=0) - return self._gather_zero_prefixed_concepts(concept_states_with_zero, prefill_metadata) - - def _repeat_shift_source_states_packed(self, - concept_states_with_zero: torch.Tensor, - prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: - """Gather zero-prefixed packed concept source states to token rows.""" - return self._gather_zero_prefixed_concepts(concept_states_with_zero, prefill_metadata) - def _build_encoder_concept_states_packed(self, encoder_raw_states: list[torch.Tensor], prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: """Build packed chunk-level encoder states used by the concept predictor.""" - chunks = [self._merge_chunks_packed(state, prefill_metadata) for state in encoder_raw_states[:-1]] + chunks = [self.concept_ops.merge_chunks_packed(state, prefill_metadata) for state in encoder_raw_states[:-1]] states = torch.stack(chunks, dim=-2) return self.concept_predictor.normalize_encoder_concept_states(states) @@ -991,10 +495,10 @@ def _build_prefill_concept_request(self, prefill_metadata: ConceptPrefillMetadata, concept_metadata: ConceptMetadata) -> _ConceptPredictorRequest: """Build concept predictor inputs for compact packed prefill.""" - concept_hidden = self._merge_chunks_packed(hidden_states, prefill_metadata) + concept_hidden = self.concept_ops.merge_chunks_packed(hidden_states, prefill_metadata) concept_hidden = self.concept_vq_input_norm(concept_hidden) encoder_concept_states = self._build_encoder_concept_states_packed(encoder_raw_states, prefill_metadata) - concept_attn_metadata = self._build_concept_prefill_metadata( + concept_attn_metadata = self.concept_ops.build_concept_prefill_metadata( concept_metadata.attn_metadata, prefill_metadata, ) @@ -1046,39 +550,10 @@ def _prepare_decoder_route_sources(self, zero_chunk = torch.zeros_like(concept_states[:1]) concept_states = torch.cat((zero_chunk, concept_states), dim=0) concept_states = self.decoder_read_concept_shared_source_norm(concept_states) - concept_states = self._repeat_shift_source_states_packed(concept_states, decoder_concepts.prefill_metadata) + concept_states = self.concept_ops.repeat_shift_source_states_packed(concept_states, + decoder_concepts.prefill_metadata) return decoder_encoder_states, concept_states - def _build_concept_prefill_metadata(self, - token_attn_metadata: Any, - prefill_metadata: ConceptPrefillMetadata): - """Build chunk-stream attention metadata for packed prefill.""" - concept_q_seqlens = prefill_metadata.concept_q_seqlens - concept_q_start_loc = prefill_metadata.concept_q_start_loc - concept_cu_seqlens = torch.nn.functional.pad( - torch.cumsum(concept_q_seqlens, dim=0, dtype=torch.int32), - (1, 0), - ) - max_concept_seqlen = int(prefill_metadata.max_concepts_per_request) - - updates = dict( - is_decoding=False, - q_start_loc=concept_q_start_loc, - q_seqlens=concept_q_seqlens, - kv_seqlens=concept_q_seqlens, - cu_seqlens_q=concept_cu_seqlens, - cu_seqlens_k=concept_cu_seqlens, - ) - if hasattr(token_attn_metadata, 'kv_start_loc'): - updates['kv_start_loc'] = concept_q_start_loc - if hasattr(token_attn_metadata, 'kv_flatten_size'): - updates['kv_flatten_size'] = int(prefill_metadata.num_concepts_total) - if hasattr(token_attn_metadata, 'max_q_seqlen'): - updates['max_q_seqlen'] = max_concept_seqlen - if hasattr(token_attn_metadata, 'max_kv_seqlen'): - updates['max_kv_seqlen'] = max_concept_seqlen - return replace(token_attn_metadata, **updates) - def _forward_prefill_packed(self, hidden_states: torch.Tensor, position_ids: torch.Tensor, @@ -1088,7 +563,7 @@ def _forward_prefill_packed(self, if (concept_caches.encoder_past_key_values is None or concept_caches.concept_past_key_values is None or concept_caches.decoder_past_key_values is None): raise RuntimeError('ConceptLM prefill requires encoder, concept, and decoder KV caches.') - prefill_metadata = self._build_prefill_metadata(concept_metadata.attn_metadata, position_ids) + prefill_metadata = self.concept_ops.build_prefill_metadata(concept_metadata.attn_metadata, position_ids) encoder_output = self._encode( hidden_states, position_ids, @@ -1104,15 +579,17 @@ def _forward_prefill_packed(self, concept_metadata, ) concept_output = self._run_concept_predictor(concept_request, concept_caches) - self._write_prefill_state_caches_eager_( - concept_caches, - concept_metadata, + self.concept_ops.write_prefill_state_caches_eager( + concept_caches.chunk_source_state, + concept_caches.last_raw_states, + concept_caches.last_final_state, + concept_metadata.state_ids, prefill_metadata, self._build_decode_chunk_source_states(encoder_output), concept_output.predicted_vectors, concept_output.raw_states, ) - repeated_concepts = self._repeat_shift_packed(concept_output.predicted_vectors, prefill_metadata) + repeated_concepts = self.concept_ops.repeat_shift_packed(concept_output.predicted_vectors, prefill_metadata) decoder_concepts = _DecoderConceptInput.for_prefill( repeated_concepts, concept_output.raw_states, @@ -1148,7 +625,7 @@ def _forward_decode(self, raise RuntimeError('ConceptLM decode requires cached last concept states.') hidden_states, decode_position_ids = self._normalize_decode_inputs(hidden_states, position_ids) - decode_metadata = self._build_decode_metadata( + decode_metadata = self.concept_ops.build_decode_metadata( decode_position_ids, concept_metadata.state_ids, hidden_states.size(0), diff --git a/lmdeploy/pytorch/nn/conceptlm.py b/lmdeploy/pytorch/nn/conceptlm.py index d36e037737..9ce2c60cff 100644 --- a/lmdeploy/pytorch/nn/conceptlm.py +++ b/lmdeploy/pytorch/nn/conceptlm.py @@ -1,7 +1,11 @@ # Copyright (c) OpenMMLab. All rights reserved. +from typing import Any + +import torch from torch import Tensor, nn from lmdeploy.pytorch.backends import OpType, get_backend +from lmdeploy.pytorch.backends.conceptlm import ConceptDecodeMetadata, ConceptPrefillMetadata class ConceptLMRuntimeOps(nn.Module): @@ -11,11 +15,129 @@ class ConceptLMRuntimeOps(nn.Module): Triton kernel launchers. """ - def __init__(self): + def __init__(self, config): super().__init__() backend = get_backend() builder = backend.get_layer_impl_builder(OpType.ConceptLMRuntimeOps) - self.impl = builder.build() + self.impl = builder.build(config) + + def flatten_decode_position_ids(self, position_ids: Tensor, batch_size: int, device: torch.device) -> Tensor: + """Normalize decode position ids to one absolute position per batch + row.""" + return self.impl.flatten_decode_position_ids(position_ids, batch_size, device) + + def build_decode_metadata(self, + position_ids: Tensor, + state_ids: Tensor | None, + batch_size: int, + device: torch.device) -> ConceptDecodeMetadata: + """Build fixed-shape decode metadata from engine state ids.""" + return self.impl.build_decode_metadata(position_ids, state_ids, batch_size, device) + + def select_decode_last_state_rows(self, + last_state: Tensor | None, + last_final_state: Tensor, + last_raw_states: Tensor, + decode_metadata: ConceptDecodeMetadata) -> tuple[Tensor, Tensor]: + """Gather latest concept state rows for decode.""" + return self.impl.select_decode_last_state_rows( + last_state, + last_final_state, + last_raw_states, + decode_metadata, + ) + + def decode_concept_read_mask(self, decode_metadata: ConceptDecodeMetadata) -> Tensor: + """Return rows whose current decode token should read a cached + concept.""" + return self.impl.decode_concept_read_mask(decode_metadata) + + def build_concept_decode_metadata_static(self, token_attn_metadata: Any, + decode_metadata: ConceptDecodeMetadata): + """Build fixed-shape concept-stream decode metadata.""" + return self.impl.build_concept_decode_metadata_static(token_attn_metadata, decode_metadata) + + def build_prefill_metadata(self, token_attn_metadata: Any, position_ids: Tensor) -> ConceptPrefillMetadata: + """Build packed token-to-concept metadata for batched prefill.""" + return self.impl.build_prefill_metadata(token_attn_metadata, position_ids) + + def merge_chunks_packed(self, hidden_states: Tensor, prefill_metadata: ConceptPrefillMetadata) -> Tensor: + """Merge packed token states into compact concept rows.""" + return self.impl.merge_chunks_packed(hidden_states, prefill_metadata) + + def repeat_shift_packed(self, concept_states: Tensor, prefill_metadata: ConceptPrefillMetadata) -> Tensor: + """Gather compact concept states back to packed token rows.""" + return self.impl.repeat_shift_packed(concept_states, prefill_metadata) + + def repeat_shift_source_states_packed(self, concept_states_with_zero: Tensor, + prefill_metadata: ConceptPrefillMetadata) -> Tensor: + """Gather zero-prefixed compact concept route states to token rows.""" + return self.impl.repeat_shift_source_states_packed(concept_states_with_zero, prefill_metadata) + + def build_concept_prefill_metadata(self, token_attn_metadata: Any, prefill_metadata: ConceptPrefillMetadata): + """Build chunk-stream attention metadata for packed prefill.""" + return self.impl.build_concept_prefill_metadata(token_attn_metadata, prefill_metadata) + + def write_prefill_state_caches_eager( + self, + chunk_source_state: Tensor | None, + last_raw_states: Tensor | None, + last_final_state: Tensor | None, + state_ids: Tensor | None, + prefill_metadata: ConceptPrefillMetadata, + source_states: Tensor, + predicted_vectors: Tensor, + concept_raw_states: list[Tensor], + ) -> None: + """Seed decode state caches from a completed prefill forward.""" + return self.impl.write_prefill_state_caches_eager( + chunk_source_state, + last_raw_states, + last_final_state, + state_ids, + prefill_metadata, + source_states, + predicted_vectors, + concept_raw_states, + ) + + def stack_concept_raw_states(self, concept_raw_states: list[Tensor]) -> Tensor: + """Stack raw concept-layer states.""" + return self.impl.stack_concept_raw_states(concept_raw_states) + + def snapshot_decode_concept_kv(self, concept_past_key_values: list[list[Tensor]], + concept_attn_metadata: Any) -> list[tuple[Tensor, Tensor]]: + """Snapshot concept KV slots that dummy rows may overwrite.""" + return self.impl.snapshot_decode_concept_kv(concept_past_key_values, concept_attn_metadata) + + def restore_decode_concept_kv(self, concept_past_key_values: list[list[Tensor]], concept_attn_metadata: Any, + saved_kv: list[tuple[Tensor, Tensor]], restore_mask: Tensor) -> None: + """Restore concept KV slots for dummy rows.""" + return self.impl.restore_decode_concept_kv( + concept_past_key_values, + concept_attn_metadata, + saved_kv, + restore_mask, + ) + + def write_decode_concept_states( + self, + last_raw_state_cache: Tensor, + last_final_state_cache: Tensor, + predicted_vectors: Tensor, + raw_states: list[Tensor], + state_ids: Tensor, + update_mask: Tensor, + ) -> None: + """Write newly emitted decode concept states.""" + return self.impl.write_decode_concept_states( + last_raw_state_cache, + last_final_state_cache, + predicted_vectors, + raw_states, + state_ids, + update_mask, + ) def decode_chunk_state_update( self, From fd4da70cb2454e8f079e58367a08455babd270cc Mon Sep 17 00:00:00 2001 From: grimoire Date: Sun, 26 Jul 2026 17:01:01 +0800 Subject: [PATCH 08/16] Propagate max query length to CUDA attention --- lmdeploy/pytorch/backends/cuda/attention/fa3.py | 2 ++ lmdeploy/pytorch/backends/cuda/op_backend.py | 1 + 2 files changed, 3 insertions(+) diff --git a/lmdeploy/pytorch/backends/cuda/attention/fa3.py b/lmdeploy/pytorch/backends/cuda/attention/fa3.py index 84aeefa189..07ff18d9f8 100644 --- a/lmdeploy/pytorch/backends/cuda/attention/fa3.py +++ b/lmdeploy/pytorch/backends/cuda/attention/fa3.py @@ -58,6 +58,8 @@ def _get_max_q_seqlen( attn_metadata: TritonAttentionMetadata, ) -> int: """Get max q seqlen.""" + if attn_metadata.max_q_seqlen is not None: + return attn_metadata.max_q_seqlen max_q_seqlen = query.numel() // (query.size(-1) * query.size(-2)) if attn_metadata.is_decoding: batch_size = attn_metadata.q_seqlens.size(0) diff --git a/lmdeploy/pytorch/backends/cuda/op_backend.py b/lmdeploy/pytorch/backends/cuda/op_backend.py index 1464472cba..447cde2ee3 100644 --- a/lmdeploy/pytorch/backends/cuda/op_backend.py +++ b/lmdeploy/pytorch/backends/cuda/op_backend.py @@ -252,6 +252,7 @@ def update_step_context(cls, step_context): cu_seqlens_q=cu_seqlens_q, cu_seqlens_k=cu_seqlens_k, max_kv_seqlen=step_context.max_kv_seqlen, + max_q_seqlen=step_context.max_q_seqlen, ) if step_context.is_decoding: if use_flash_mla: From 954b45f93ab952007d9420cea74bdd3aa2fdcdff Mon Sep 17 00:00:00 2001 From: grimoire Date: Sun, 26 Jul 2026 18:03:31 +0800 Subject: [PATCH 09/16] Add ConceptLM runtime support --- lmdeploy/pytorch/backends/conceptlm.py | 658 ++------------- lmdeploy/pytorch/backends/cuda/conceptlm.py | 81 +- .../pytorch/backends/default/conceptlm.py | 768 +++++++++++++++++- lmdeploy/pytorch/kernels/cuda/conceptlm.py | 274 +++++++ .../pytorch/models/intern_ncp/modeling.py | 423 ++-------- lmdeploy/pytorch/nn/conceptlm.py | 216 ++--- 6 files changed, 1335 insertions(+), 1085 deletions(-) diff --git a/lmdeploy/pytorch/backends/conceptlm.py b/lmdeploy/pytorch/backends/conceptlm.py index 5a819e6eec..0e88d40743 100644 --- a/lmdeploy/pytorch/backends/conceptlm.py +++ b/lmdeploy/pytorch/backends/conceptlm.py @@ -1,20 +1,63 @@ # Copyright (c) OpenMMLab. All rights reserved. from abc import ABC, abstractmethod -from dataclasses import dataclass, replace +from dataclasses import dataclass from typing import Any import torch from torch import Tensor -from torch.nn import functional as F from transformers.configuration_utils import PretrainedConfig @dataclass -class ConceptChunkStateUpdateResult: - """Fixed-shape result of one decode chunk-source state update.""" +class ConceptChunkInput: + """Concept-stream input rows prepared from token-stream encoder states. - concept_input_states: Tensor - concept_update_mask: Tensor + ``source_states`` has a unified layout for prefill and decode: + ``[concept_rows, num_sources, hidden]``. Source row 0 is the hidden state + consumed by the concept predictor; rows 1: are encoder states consumed by + ConceptRoute/SelfDD. + """ + + source_states: Tensor + position_ids: Tensor + attn_metadata: Any + state_ids: Tensor | None = None + update_mask: Tensor | None = None + prefill_metadata: 'ConceptPrefillMetadata | None' = None + decode_metadata: 'ConceptDecodeMetadata | None' = None + + @property + def is_decoding(self) -> bool: + """Whether this input came from the fixed-shape decode path.""" + return self.decode_metadata is not None + + +@dataclass +class ConceptForwardContext: + """Temporary state needed around one concept-predictor forward.""" + + saved_kv: list[tuple[Tensor, Tensor]] | None = None + previous_final_state: Tensor | None = None + previous_raw_states: Tensor | None = None + + +@dataclass +class ConceptDecoderInput: + """Concept states consumed by the token decoder stack.""" + + final_state: Tensor + route_states: Tensor + + +@dataclass +class ConceptRuntimeCaches: + """Backend-facing ConceptLM runtime cache views.""" + + chunk_source_state: Tensor | None = None + last_state: Tensor | None = None + last_raw_states: Tensor | None = None + last_final_state: Tensor | None = None + concept_past_key_values: list[list[Tensor]] | None = None @dataclass @@ -41,6 +84,7 @@ class ConceptPrefillMetadata: concept_q_start_loc: Tensor concept_position_ids: Tensor merge_token_to_concept: Tensor + merge_token_start_ids: Tensor merge_token_counts: Tensor merge_first_token_ids: Tensor merge_last_token_ids: Tensor @@ -51,65 +95,8 @@ class ConceptPrefillMetadata: max_concepts_per_request: int -@dataclass -class _PrefillTokenLayout: - """Packed token-stream layout derived from prefill attention metadata.""" - - q_seqlens: Tensor - q_start_loc: Tensor - q_seqlens_long: Tensor - q_start_loc_long: Tensor - token_seq: Tensor - token_pos: Tensor - total_tokens: int - - -@dataclass -class _PrefillConceptLayout: - """Compact chunk-token stream layout used by concept predictor prefill.""" - - q_seqlens: Tensor - q_seqlens_long: Tensor - q_start_loc: Tensor - q_start_loc_long: Tensor - seq: Tensor - local_ids: Tensor - position_ids: Tensor - num_total: int - max_per_request: int - - -@dataclass -class _PrefillMergeLayout: - """Token-to-concept merge metadata for compact prefill.""" - - token_to_concept: Tensor - token_counts: Tensor - first_token_ids: Tensor - last_token_ids: Tensor - short_concept_mask: Tensor - - -def _flatten_decode_position_ids(position_ids: Tensor, batch_size: int, device: torch.device) -> Tensor: - """Normalize decode position ids to one absolute position per batch row.""" - if position_ids.dim() == 0: - position_ids = position_ids.view(1) - if position_ids.dim() == 1: - return position_ids.to(device=device, dtype=torch.long) - position_ids = position_ids.reshape(-1) - if position_ids.numel() == batch_size: - return position_ids.to(device=device, dtype=torch.long) - assert position_ids.numel() % batch_size == 0, ( - f'Cannot map position_ids with {position_ids.numel()} elements to batch size {batch_size}.') - return position_ids.reshape(-1, batch_size)[-1].to(device=device, dtype=torch.long) - - class ConceptLMRuntimeOpsImpl(ABC): - """ConceptLM runtime operation implementation. - - Model-specific runtime/cache operations live behind this single backend interface. That keeps model code free from - direct kernel calls while avoiding one OpType/nn module per small ConceptLM state operation. - """ + """Backend contract for ConceptLM runtime/cache operations.""" def __init__(self, config: PretrainedConfig): self.config = config @@ -117,542 +104,53 @@ def __init__(self, config: PretrainedConfig): self.merge_method = getattr(config, 'concept_chunk_merge_method', 'meanpooling') self.shift_feature = bool(getattr(config, 'concept_shift_feature', True)) - @staticmethod - def concept_count_from_seq_len(seq_len: int, chunk_size: int) -> int: - """Return reference ConceptLM chunk count for one request length.""" - seq_len = int(seq_len or 0) - if seq_len <= 0: - return 0 - if seq_len < chunk_size: - return 1 - return seq_len // chunk_size - - @staticmethod - def concept_counts_from_q_seqlens(q_seqlens: Tensor, chunk_size: int) -> Tensor: - """Vectorized version of ``concept_count_from_seq_len``.""" - counts = torch.div(q_seqlens, chunk_size, rounding_mode='floor').clamp(min=1) - return torch.where(q_seqlens > 0, counts, torch.zeros_like(q_seqlens)) - - @staticmethod - def repeat_slot_ids(token_pos: Tensor, chunk_size: int, shift_feature: bool) -> Tensor: - """Return local concept slot read by each token after shift - semantics.""" - if shift_feature: - return torch.div(token_pos + 1, chunk_size, rounding_mode='floor') - 1 - return torch.div(token_pos, chunk_size, rounding_mode='floor') - 1 - + @abstractmethod def flatten_decode_position_ids(self, position_ids: Tensor, batch_size: int, device: torch.device) -> Tensor: """Normalize decode position ids to one absolute position per batch row.""" - return _flatten_decode_position_ids(position_ids, batch_size, device) - - def build_decode_metadata(self, - position_ids: Tensor, - state_ids: Tensor | None, - batch_size: int, - device: torch.device) -> ConceptDecodeMetadata: - """Build fixed-shape decode metadata from engine state ids.""" - if state_ids is None: - raise RuntimeError('ConceptLM decode requires state_ids.') - state_ids = state_ids.to(device=device, dtype=torch.long).reshape(-1) - if state_ids.numel() != batch_size: - raise ValueError(f'Expected {batch_size} decode state ids, got {state_ids.numel()}.') - valid_state_mask = state_ids >= 0 - return ConceptDecodeMetadata( - position_ids=position_ids, - state_ids=state_ids, - safe_state_ids=state_ids.clamp(min=0), - valid_state_mask=valid_state_mask, - ) - - @staticmethod - def select_decode_state_rows(state_cache: Tensor, decode_metadata: ConceptDecodeMetadata) -> Tensor: - """Gather state-cache rows and zero out padded decode rows.""" - rows = state_cache.index_select(0, decode_metadata.safe_state_ids) - mask_shape = (decode_metadata.valid_state_mask.size(0), ) + (1, ) * (rows.dim() - 1) - valid_mask = decode_metadata.valid_state_mask.view(mask_shape) - return torch.where(valid_mask, rows, torch.zeros_like(rows)) - - def select_decode_last_state_rows(self, - last_state: Tensor | None, - last_final_state: Tensor, - last_raw_states: Tensor, - decode_metadata: ConceptDecodeMetadata) -> tuple[Tensor, Tensor]: - """Gather packed last-concept state rows once and return final/raw - views.""" - if last_state is not None: - rows = self.select_decode_state_rows(last_state, decode_metadata) - return rows[:, 0], rows[:, 1:] - - return ( - self.select_decode_state_rows(last_final_state, decode_metadata), - self.select_decode_state_rows(last_raw_states, decode_metadata), - ) - - def decode_concept_read_mask(self, decode_metadata: ConceptDecodeMetadata) -> Tensor: - """Return rows whose current decode token should read a cached - concept.""" - repeat_slots = self.repeat_slot_ids( - decode_metadata.position_ids, - self.chunk_size, - self.shift_feature, - ) - return decode_metadata.valid_state_mask & (repeat_slots >= 0) - - def build_concept_decode_metadata_static(self, token_attn_metadata: Any, - decode_metadata: ConceptDecodeMetadata): - """Build fixed-shape concept-stream decode metadata.""" - device = decode_metadata.position_ids.device - batch_size = decode_metadata.position_ids.numel() - q_seqlens = token_attn_metadata.q_seqlens - q_start_loc = token_attn_metadata.q_start_loc - kv_seqlens = token_attn_metadata.kv_seqlens - q_dtype = q_seqlens.dtype - q_start_dtype = q_start_loc.dtype - kv_dtype = kv_seqlens.dtype - - concept_q_seqlens = torch.ones((batch_size, ), dtype=q_dtype, device=device) - concept_q_start_loc = torch.arange(batch_size, dtype=q_start_dtype, device=device) - concept_cu_seqlens = F.pad(torch.cumsum(concept_q_seqlens, dim=0, dtype=torch.int32), (1, 0)) - concept_kv_seqlens = torch.div( - decode_metadata.position_ids + 1, - self.chunk_size, - rounding_mode='floor', - ).clamp(min=1).to(dtype=kv_dtype) - - updates = dict( - is_decoding=True, - block_offsets=token_attn_metadata.block_offsets, - q_start_loc=concept_q_start_loc, - q_seqlens=concept_q_seqlens, - kv_seqlens=concept_kv_seqlens, - cu_seqlens_q=concept_cu_seqlens, - cu_seqlens_k=concept_cu_seqlens, - ) - if hasattr(token_attn_metadata, 'kv_start_loc'): - updates['kv_start_loc'] = concept_kv_seqlens - concept_q_seqlens.to(dtype=concept_kv_seqlens.dtype) - if hasattr(token_attn_metadata, 'kv_flatten_size'): - updates['kv_flatten_size'] = batch_size - if hasattr(token_attn_metadata, 'max_q_seqlen'): - updates['max_q_seqlen'] = 1 - if hasattr(token_attn_metadata, 'max_kv_seqlen'): - updates['max_kv_seqlen'] = getattr(token_attn_metadata, 'max_kv_seqlen') - if hasattr(token_attn_metadata, 'scheduler_metadata'): - updates['scheduler_metadata'] = None - if hasattr(token_attn_metadata, 'tile_scheduler_metadata'): - updates['tile_scheduler_metadata'] = None - if hasattr(token_attn_metadata, 'num_splits'): - updates['num_splits'] = None - if hasattr(token_attn_metadata, 'fill_seqlens'): - updates['fill_seqlens'] = concept_q_seqlens - return replace(token_attn_metadata, **updates) - - def _get_max_concepts_per_request(self, token_attn_metadata: Any, concept_q_seqlens: Tensor) -> int: - """Return per-request concept attention bound without hidden context - access.""" - max_q_seqlen = getattr(token_attn_metadata, 'max_q_seqlen', None) - if max_q_seqlen is not None: - return self.concept_count_from_seq_len(int(max_q_seqlen), self.chunk_size) - - # Test/direct-call fallback. Serving should use the scheduler-provided - # Python max_q_seqlen above, as in DSV4 metadata construction. - return int(concept_q_seqlens.max().item()) if concept_q_seqlens.numel() > 0 else 0 - - @staticmethod - def _build_prefill_token_layout(token_attn_metadata: Any, position_ids: Tensor) -> _PrefillTokenLayout: - """Build packed token-stream positions from prefill attention - metadata.""" - q_seqlens = token_attn_metadata.q_seqlens - q_start_loc = getattr(token_attn_metadata, 'q_start_loc', None) - if q_start_loc is None: - q_start_loc = F.pad(torch.cumsum(q_seqlens, dim=0, dtype=torch.int32), (1, 0))[:-1] - - total_tokens = int(position_ids.numel()) - - q_seqlens_long = q_seqlens.to(dtype=torch.long, device=position_ids.device) - q_start_loc_long = q_start_loc.to(dtype=torch.long, device=position_ids.device) - cu_q_seqlens = getattr(token_attn_metadata, 'cu_seqlens_q', None) - if cu_q_seqlens is None: - cu_q_seqlens = F.pad(torch.cumsum(q_seqlens, dim=0, dtype=torch.int32), (1, 0)) - cu_q_seqlens_long = cu_q_seqlens.to(dtype=torch.long, device=position_ids.device) - - token_ids = torch.arange(total_tokens, dtype=torch.long, device=position_ids.device) - token_seq = torch.searchsorted(cu_q_seqlens_long[1:], token_ids, right=True) - token_pos = token_ids - cu_q_seqlens_long[token_seq] - - return _PrefillTokenLayout( - q_seqlens=q_seqlens, - q_start_loc=q_start_loc, - q_seqlens_long=q_seqlens_long, - q_start_loc_long=q_start_loc_long, - token_seq=token_seq, - token_pos=token_pos, - total_tokens=total_tokens, - ) - - def _build_prefill_concept_layout(self, token_attn_metadata: Any, position_ids: Tensor, - token_layout: _PrefillTokenLayout) -> _PrefillConceptLayout: - """Build compact chunk-token positions used by concept attention.""" - concept_q_seqlens_long = self.concept_counts_from_q_seqlens(token_layout.q_seqlens_long, self.chunk_size) - concept_q_seqlens = concept_q_seqlens_long.to(dtype=token_layout.q_seqlens.dtype, - device=token_layout.q_seqlens.device) - concept_q_start_loc = F.pad(torch.cumsum(concept_q_seqlens, dim=0, dtype=torch.int32), (1, 0))[:-1] - concept_q_start_loc_long = concept_q_start_loc.to(dtype=torch.long, device=position_ids.device) - concept_cu_seqlens_long = F.pad(torch.cumsum(concept_q_seqlens_long, dim=0), (1, 0)) - - # TODO: remove this eager compact-size scalar read by moving ConceptLM - # concept-stream allocation/compaction into the engine/backend contract, - # like DSV4's precomputed metadata and Qwen3.5's state-cache metadata. - num_concepts_total = int(concept_cu_seqlens_long[-1].item()) - concept_ids = torch.arange(num_concepts_total, dtype=torch.long, device=position_ids.device) - concept_seq = torch.searchsorted(concept_cu_seqlens_long[1:], concept_ids, right=True) - local_concept_ids = concept_ids - concept_cu_seqlens_long[concept_seq] - concept_token_start = token_layout.q_start_loc_long[concept_seq] + local_concept_ids * self.chunk_size - concept_position_ids = position_ids[concept_token_start] - max_concepts_per_request = self._get_max_concepts_per_request(token_attn_metadata, concept_q_seqlens_long) - - return _PrefillConceptLayout( - q_seqlens=concept_q_seqlens, - q_seqlens_long=concept_q_seqlens_long, - q_start_loc=concept_q_start_loc, - q_start_loc_long=concept_q_start_loc_long, - seq=concept_seq, - local_ids=local_concept_ids, - position_ids=concept_position_ids, - num_total=num_concepts_total, - max_per_request=max_concepts_per_request, - ) - - def _build_prefill_repeat_ids(self, token_layout: _PrefillTokenLayout, - concept_layout: _PrefillConceptLayout) -> Tensor: - """Map packed token rows to the concept row visible after shift.""" - seq_concept_start = concept_layout.q_start_loc_long[token_layout.token_seq] - seq_concept_count = concept_layout.q_seqlens_long[token_layout.token_seq] - repeat_slots = self.repeat_slot_ids(token_layout.token_pos, self.chunk_size, self.shift_feature) - valid_repeat = (repeat_slots >= 0) & (repeat_slots < seq_concept_count) - return torch.where( - valid_repeat, - seq_concept_start + repeat_slots, - torch.full_like(repeat_slots, -1), - ) - - def _build_prefill_merge_layout(self, token_layout: _PrefillTokenLayout, - concept_layout: _PrefillConceptLayout) -> _PrefillMergeLayout: - """Build token-to-concept merge metadata for mean/first/last modes.""" - seq_concept_start = concept_layout.q_start_loc_long[token_layout.token_seq] - seq_concept_count = concept_layout.q_seqlens_long[token_layout.token_seq] - merge_slots = torch.div(token_layout.token_pos, self.chunk_size, rounding_mode='floor') - valid_merge = (merge_slots >= 0) & (merge_slots < seq_concept_count) - merge_token_to_concept = torch.where( - valid_merge, - seq_concept_start + merge_slots, - torch.full_like(merge_slots, -1), - ) - safe_merge_ids = merge_token_to_concept.clamp(min=0) - merge_token_counts = torch.zeros(concept_layout.num_total, - dtype=torch.int32, - device=token_layout.token_pos.device) - merge_token_counts.index_add_(0, safe_merge_ids, valid_merge.to(dtype=torch.int32)) - - concept_seq_len = token_layout.q_seqlens_long[concept_layout.seq] - merge_first_pos = torch.where( - concept_seq_len < self.chunk_size, - torch.zeros_like(concept_layout.local_ids), - concept_layout.local_ids * self.chunk_size, - ) - merge_last_pos = torch.where( - concept_seq_len < self.chunk_size, - (concept_seq_len - 1).clamp(min=0), - concept_layout.local_ids * self.chunk_size + self.chunk_size - 1, - ) - merge_first_token_ids = token_layout.q_start_loc_long[concept_layout.seq] + merge_first_pos - merge_last_token_ids = token_layout.q_start_loc_long[concept_layout.seq] + merge_last_pos - merge_short_concept_mask = concept_seq_len < self.chunk_size - - return _PrefillMergeLayout( - token_to_concept=merge_token_to_concept, - token_counts=merge_token_counts, - first_token_ids=merge_first_token_ids, - last_token_ids=merge_last_token_ids, - short_concept_mask=merge_short_concept_mask, - ) - - def build_prefill_metadata(self, token_attn_metadata: Any, position_ids: Tensor) -> ConceptPrefillMetadata: - """Build packed token-to-concept metadata for batched prefill.""" - token_layout = self._build_prefill_token_layout(token_attn_metadata, position_ids) - concept_layout = self._build_prefill_concept_layout(token_attn_metadata, position_ids, token_layout) - token_to_concept = self._build_prefill_repeat_ids(token_layout, concept_layout) - merge_layout = self._build_prefill_merge_layout(token_layout, concept_layout) - - return ConceptPrefillMetadata( - token_q_seqlens=token_layout.q_seqlens, - token_q_start_loc=token_layout.q_start_loc, - concept_q_seqlens=concept_layout.q_seqlens, - concept_q_start_loc=concept_layout.q_start_loc, - concept_position_ids=concept_layout.position_ids, - merge_token_to_concept=merge_layout.token_to_concept, - merge_token_counts=merge_layout.token_counts, - merge_first_token_ids=merge_layout.first_token_ids, - merge_last_token_ids=merge_layout.last_token_ids, - merge_short_concept_mask=merge_layout.short_concept_mask, - token_to_concept=token_to_concept, - num_tokens_total=token_layout.total_tokens, - num_concepts_total=concept_layout.num_total, - max_concepts_per_request=concept_layout.max_per_request, - ) - - @staticmethod - def _merge_chunks_mean_packed(hidden_states: Tensor, prefill_metadata: ConceptPrefillMetadata) -> Tensor: - """Mean-pool packed token states by precomputed concept ids.""" - merge_token_to_concept = prefill_metadata.merge_token_to_concept.to(device=hidden_states.device) - valid_merge = (merge_token_to_concept >= 0).to(dtype=hidden_states.dtype).unsqueeze(-1) - safe_merge_ids = merge_token_to_concept.clamp(min=0) - merged = hidden_states.new_zeros((prefill_metadata.num_concepts_total, hidden_states.size(-1))) - merged.index_add_(0, safe_merge_ids, hidden_states * valid_merge) - counts = prefill_metadata.merge_token_counts.clamp(min=1).to(device=hidden_states.device, - dtype=hidden_states.dtype) - return merged / counts.unsqueeze(-1) - - def merge_chunks_packed(self, hidden_states: Tensor, prefill_metadata: ConceptPrefillMetadata) -> Tensor: - """Merge packed token states into packed per-request concept states.""" - if self.merge_method == 'first': - merged = hidden_states[prefill_metadata.merge_first_token_ids] - short_mean = self._merge_chunks_mean_packed(hidden_states, prefill_metadata) - return torch.where(prefill_metadata.merge_short_concept_mask[:, None], short_mean, merged) - if self.merge_method == 'last': - merged = hidden_states[prefill_metadata.merge_last_token_ids] - short_mean = self._merge_chunks_mean_packed(hidden_states, prefill_metadata) - return torch.where(prefill_metadata.merge_short_concept_mask[:, None], short_mean, merged) - - return self._merge_chunks_mean_packed(hidden_states, prefill_metadata) - - @staticmethod - def _gather_zero_prefixed_concepts(concept_states_with_zero: Tensor, - prefill_metadata: ConceptPrefillMetadata) -> Tensor: - """Gather zero-prefixed concept rows to packed token rows.""" - token_to_concept = prefill_metadata.token_to_concept.to(device=concept_states_with_zero.device) - gather_ids = torch.clamp(token_to_concept + 1, min=0) - return concept_states_with_zero[gather_ids] - - def repeat_shift_packed(self, concept_states: Tensor, prefill_metadata: ConceptPrefillMetadata) -> Tensor: - """Gather packed concept states back to packed token states.""" - concept_states_with_zero = torch.cat((torch.zeros_like(concept_states[:1]), concept_states), dim=0) - return self._gather_zero_prefixed_concepts(concept_states_with_zero, prefill_metadata) - - def repeat_shift_source_states_packed(self, concept_states_with_zero: Tensor, - prefill_metadata: ConceptPrefillMetadata) -> Tensor: - """Gather zero-prefixed packed concept source states to token rows.""" - return self._gather_zero_prefixed_concepts(concept_states_with_zero, prefill_metadata) - - def build_concept_prefill_metadata(self, token_attn_metadata: Any, prefill_metadata: ConceptPrefillMetadata): - """Build chunk-stream attention metadata for packed prefill.""" - concept_q_seqlens = prefill_metadata.concept_q_seqlens - concept_q_start_loc = prefill_metadata.concept_q_start_loc - concept_cu_seqlens = F.pad( - torch.cumsum(concept_q_seqlens, dim=0, dtype=torch.int32), - (1, 0), - ) - max_concept_seqlen = int(prefill_metadata.max_concepts_per_request) - - updates = dict( - is_decoding=False, - q_start_loc=concept_q_start_loc, - q_seqlens=concept_q_seqlens, - kv_seqlens=concept_q_seqlens, - cu_seqlens_q=concept_cu_seqlens, - cu_seqlens_k=concept_cu_seqlens, - ) - if hasattr(token_attn_metadata, 'kv_start_loc'): - updates['kv_start_loc'] = concept_q_start_loc - if hasattr(token_attn_metadata, 'kv_flatten_size'): - updates['kv_flatten_size'] = int(prefill_metadata.num_concepts_total) - if hasattr(token_attn_metadata, 'max_q_seqlen'): - updates['max_q_seqlen'] = max_concept_seqlen - if hasattr(token_attn_metadata, 'max_kv_seqlen'): - updates['max_kv_seqlen'] = max_concept_seqlen - return replace(token_attn_metadata, **updates) - - def merge_prefill_tail_chunk_states(self, source_states: Tensor, - prefill_metadata: ConceptPrefillMetadata) -> Tensor: - """Build per-request partial chunk accumulator rows after prefill.""" - device = source_states.device - q_seqlens = prefill_metadata.token_q_seqlens.to(device=device, dtype=torch.long) - q_start_loc = prefill_metadata.token_q_start_loc.to(device=device, dtype=torch.long) - batch_size = q_seqlens.size(0) - tail_lens = torch.remainder(q_seqlens, self.chunk_size) - tail_lens = torch.where(q_seqlens < self.chunk_size, q_seqlens, tail_lens) - tail_lens = torch.where(q_seqlens > 0, tail_lens, torch.zeros_like(tail_lens)) - has_tail = tail_lens > 0 - - tail_rows = source_states.new_zeros((batch_size, source_states.size(1), source_states.size(2))) - if source_states.size(0) == 0: - return tail_rows - if self.merge_method == 'first': - token_ids = q_start_loc + q_seqlens - tail_lens - token_ids = token_ids.clamp(min=0, max=max(source_states.size(0) - 1, 0)) - rows = source_states[token_ids] - return torch.where(has_tail.view(batch_size, 1, 1), rows, tail_rows) - if self.merge_method == 'last': - token_ids = q_start_loc + q_seqlens - 1 - token_ids = token_ids.clamp(min=0, max=max(source_states.size(0) - 1, 0)) - rows = source_states[token_ids] - return torch.where(has_tail.view(batch_size, 1, 1), rows, tail_rows) - - token_ids = torch.arange(prefill_metadata.num_tokens_total, dtype=torch.long, device=device) - cu_q_seqlens = F.pad(torch.cumsum(q_seqlens, dim=0), (1, 0)) - token_seq = torch.searchsorted(cu_q_seqlens[1:], token_ids, right=True) - token_pos = token_ids - cu_q_seqlens[token_seq] - token_tail_start = q_seqlens[token_seq] - tail_lens[token_seq] - valid_tail = (tail_lens[token_seq] > 0) & (token_pos >= token_tail_start) - weighted_source = source_states * valid_tail.to(dtype=source_states.dtype).view(-1, 1, 1) - tail_rows.index_add_(0, token_seq, weighted_source) - return tail_rows - - @staticmethod - def stack_concept_raw_states(concept_raw_states: list[Tensor]) -> Tensor: - """Stack raw concept-layer states to ``[rows, concept_layers, - hidden]``.""" - return torch.stack(tuple(concept_raw_states), dim=1) - - def write_prefill_state_caches_eager( - self, - chunk_source_state: Tensor | None, - last_raw_states: Tensor | None, - last_final_state: Tensor | None, - state_ids: Tensor | None, - prefill_metadata: ConceptPrefillMetadata, - source_states: Tensor, - predicted_vectors: Tensor, - concept_raw_states: list[Tensor], - ) -> None: - """Seed decode state caches from a completed prefill forward.""" - if state_ids is None: - return - if chunk_source_state is None or last_raw_states is None or last_final_state is None: - return - - state_ids = state_ids.to(device=source_states.device, dtype=torch.long).reshape(-1) - valid_indices = torch.nonzero(state_ids >= 0, as_tuple=False).flatten() - if valid_indices.numel() == 0: - return - - valid_state_ids = state_ids.index_select(0, valid_indices) - tail_rows = self.merge_prefill_tail_chunk_states(source_states, prefill_metadata) - chunk_source_state.index_copy_(0, valid_state_ids, - tail_rows.index_select(0, valid_indices).to(dtype=chunk_source_state.dtype)) - - concept_counts = prefill_metadata.concept_q_seqlens.to(device=source_states.device, dtype=torch.long) - concept_start = prefill_metadata.concept_q_start_loc.to(device=source_states.device, dtype=torch.long) - concept_valid_mask = (state_ids >= 0) & (concept_counts > 0) - concept_indices = torch.nonzero(concept_valid_mask, as_tuple=False).flatten() - if concept_indices.numel() == 0: - return - - concept_state_ids = state_ids.index_select(0, concept_indices) - last_concept_ids = concept_start.index_select(0, concept_indices) + concept_counts.index_select( - 0, concept_indices) - 1 - last_final_rows = predicted_vectors.index_select(0, last_concept_ids) - last_final_state.index_copy_(0, concept_state_ids, last_final_rows.to(dtype=last_final_state.dtype)) - raw_rows = self.stack_concept_raw_states(concept_raw_states).index_select(0, last_concept_ids) - last_raw_states.index_copy_(0, concept_state_ids, raw_rows.to(dtype=last_raw_states.dtype)) - - def snapshot_decode_concept_kv(self, concept_past_key_values: list[list[Tensor]], - concept_attn_metadata: Any) -> list[tuple[Tensor, Tensor]]: - """Snapshot concept KV slots that dummy non-boundary rows may - overwrite.""" - return [ - self.decode_kv_cache_snapshot( - k_cache, - v_cache, - concept_attn_metadata.block_offsets, - concept_attn_metadata.kv_seqlens, - ) - for k_cache, v_cache in concept_past_key_values - ] - - def restore_decode_concept_kv(self, concept_past_key_values: list[list[Tensor]], concept_attn_metadata: Any, - saved_kv: list[tuple[Tensor, Tensor]], restore_mask: Tensor) -> None: - """Restore concept KV slots for non-boundary and padded rows.""" - for (k_cache, v_cache), (saved_k, saved_v) in zip(concept_past_key_values, saved_kv): - self.decode_kv_cache_restore( - k_cache, - v_cache, - saved_k, - saved_v, - concept_attn_metadata.block_offsets, - concept_attn_metadata.kv_seqlens, - restore_mask, - ) - - def write_decode_concept_states( - self, - last_raw_state_cache: Tensor, - last_final_state_cache: Tensor, - predicted_vectors: Tensor, - raw_states: list[Tensor], - state_ids: Tensor, - update_mask: Tensor, - ) -> None: - """Write newly emitted concept states to persistent decode caches.""" - raw_rows = self.stack_concept_raw_states(raw_states) - self.decode_concept_state_update( - last_raw_state_cache, - last_final_state_cache, - predicted_vectors, - raw_rows, - state_ids, - update_mask, - ) + raise NotImplementedError('Not implemented.') @abstractmethod - def decode_chunk_state_update( + def build_concept_chunk_input( self, - chunk_source_state_cache: Tensor, - current_source_states: Tensor, - state_ids: Tensor, + source_states: Tensor, + token_attn_metadata: Any, position_ids: Tensor, - chunk_size: int, - merge_method: str, - ) -> tuple[Tensor, Tensor]: - """Update state cache and return concept inputs plus update mask.""" + state_ids: Tensor | None = None, + chunk_source_state_cache: Tensor | None = None, + ) -> ConceptChunkInput: + """Build concept-predictor source rows for prefill or decode.""" raise NotImplementedError('Not implemented.') @abstractmethod - def decode_kv_cache_snapshot( - self, - k_cache: Tensor, - v_cache: Tensor, - block_offsets: Tensor, - kv_seqlens: Tensor, - ) -> tuple[Tensor, Tensor]: - """Snapshot one decode KV slot per batch row.""" + def begin_concept_forward(self, chunk_input: ConceptChunkInput, + runtime_caches: ConceptRuntimeCaches) -> ConceptForwardContext: + """Prepare transient state before the concept predictor forward.""" raise NotImplementedError('Not implemented.') @abstractmethod - def decode_kv_cache_restore( + def end_concept_forward( self, - k_cache: Tensor, - v_cache: Tensor, - saved_k: Tensor, - saved_v: Tensor, - block_offsets: Tensor, - kv_seqlens: Tensor, - restore_mask: Tensor, + chunk_input: ConceptChunkInput, + runtime_caches: ConceptRuntimeCaches, + forward_context: ConceptForwardContext, + source_states: Tensor, + predicted_vectors: Tensor, + concept_raw_states: list[Tensor], ) -> None: - """Restore one decode KV slot for masked batch rows.""" + """Commit concept-predictor side effects for prefill or decode.""" raise NotImplementedError('Not implemented.') @abstractmethod - def decode_concept_state_update( + def build_decoder_concept_input( self, - last_raw_state_cache: Tensor, - last_final_state_cache: Tensor, + chunk_input: ConceptChunkInput, + runtime_caches: ConceptRuntimeCaches, + forward_context: ConceptForwardContext, predicted_vectors: Tensor, - raw_states: Tensor, - state_ids: Tensor, - update_mask: Tensor, - ) -> None: - """Write final/raw concept states for masked decode rows.""" + concept_raw_states: list[Tensor], + ) -> ConceptDecoderInput: + """Build token-decoder concept inputs for prefill or decode.""" raise NotImplementedError('Not implemented.') diff --git a/lmdeploy/pytorch/backends/cuda/conceptlm.py b/lmdeploy/pytorch/backends/cuda/conceptlm.py index 5f9c55e29c..c0749548ee 100644 --- a/lmdeploy/pytorch/backends/cuda/conceptlm.py +++ b/lmdeploy/pytorch/backends/cuda/conceptlm.py @@ -6,14 +6,93 @@ decode_concept_state_update, decode_kv_cache_restore, decode_kv_cache_snapshot, + prefill_chunk_state_update, + prefill_state_cache_update, ) from ..conceptlm import ConceptLMRuntimeOpsBuilder, ConceptLMRuntimeOpsImpl +from ..default.conceptlm import DefaultConceptLMRuntimeOpsImpl -class TritonConceptLMRuntimeOpsImpl(ConceptLMRuntimeOpsImpl): +class TritonConceptLMRuntimeOpsImpl(DefaultConceptLMRuntimeOpsImpl): """Triton implementation of ConceptLM runtime operations.""" + def merge_chunks_packed(self, hidden_states: Tensor, prefill_metadata) -> Tensor: + """Merge packed token states into packed concept rows.""" + if not hidden_states.is_cuda: + return super().merge_chunks_packed(hidden_states, prefill_metadata) + source_states = hidden_states.unsqueeze(1) + return self.prefill_chunk_state_update(source_states, prefill_metadata)[:, 0] + + def prefill_chunk_state_update(self, source_states: Tensor, prefill_metadata) -> Tensor: + """Merge prefill source states to compact concept rows.""" + if not source_states.is_cuda: + return super().prefill_chunk_state_update(source_states, prefill_metadata) + if source_states.dim() == 2: + source_states = source_states.unsqueeze(1) + return prefill_chunk_state_update( + source_states, + prefill_metadata.merge_token_start_ids, + prefill_metadata.merge_token_counts, + prefill_metadata.num_concepts_total, + self.chunk_size, + self.merge_method, + )[:, 0] + return prefill_chunk_state_update( + source_states, + prefill_metadata.merge_token_start_ids, + prefill_metadata.merge_token_counts, + prefill_metadata.num_concepts_total, + self.chunk_size, + self.merge_method, + ) + + def write_prefill_state_caches_eager( + self, + chunk_source_state: Tensor | None, + last_raw_states: Tensor | None, + last_final_state: Tensor | None, + state_ids: Tensor | None, + prefill_metadata, + source_states: Tensor, + predicted_vectors: Tensor, + concept_raw_states: list[Tensor], + ) -> None: + """Seed decode state caches from a completed CUDA prefill forward.""" + if state_ids is None: + return + if chunk_source_state is None or last_raw_states is None or last_final_state is None: + return + if not (source_states.is_cuda and chunk_source_state.is_cuda and last_raw_states.is_cuda + and last_final_state.is_cuda): + return super().write_prefill_state_caches_eager( + chunk_source_state, + last_raw_states, + last_final_state, + state_ids, + prefill_metadata, + source_states, + predicted_vectors, + concept_raw_states, + ) + + raw_rows = self.stack_concept_raw_states(concept_raw_states) + return prefill_state_cache_update( + chunk_source_state, + last_raw_states, + last_final_state, + source_states, + predicted_vectors, + raw_rows, + state_ids, + prefill_metadata.token_q_start_loc, + prefill_metadata.token_q_seqlens, + prefill_metadata.concept_q_start_loc, + prefill_metadata.concept_q_seqlens, + self.chunk_size, + self.merge_method, + ) + def decode_chunk_state_update( self, chunk_source_state_cache: Tensor, diff --git a/lmdeploy/pytorch/backends/default/conceptlm.py b/lmdeploy/pytorch/backends/default/conceptlm.py index d1c3a459bf..d307329658 100644 --- a/lmdeploy/pytorch/backends/default/conceptlm.py +++ b/lmdeploy/pytorch/backends/default/conceptlm.py @@ -1,13 +1,779 @@ # Copyright (c) OpenMMLab. All rights reserved. +from dataclasses import dataclass, replace +from typing import Any + import torch from torch import Tensor +from torch.nn import functional as F + +from ..conceptlm import ( + ConceptChunkInput, + ConceptDecodeMetadata, + ConceptDecoderInput, + ConceptForwardContext, + ConceptLMRuntimeOpsBuilder, + ConceptLMRuntimeOpsImpl, + ConceptPrefillMetadata, + ConceptRuntimeCaches, +) + + +@dataclass +class _PrefillTokenLayout: + """Packed token-stream layout derived from prefill attention metadata.""" + + q_seqlens: Tensor + q_start_loc: Tensor + q_seqlens_long: Tensor + q_start_loc_long: Tensor + token_seq: Tensor + token_pos: Tensor + total_tokens: int + + +@dataclass +class _PrefillConceptLayout: + """Compact chunk-token stream layout used by concept predictor prefill.""" -from ..conceptlm import ConceptLMRuntimeOpsBuilder, ConceptLMRuntimeOpsImpl + q_seqlens: Tensor + q_seqlens_long: Tensor + q_start_loc: Tensor + q_start_loc_long: Tensor + seq: Tensor + local_ids: Tensor + position_ids: Tensor + num_total: int + max_per_request: int + + +@dataclass +class _PrefillMergeLayout: + """Token-to-concept merge metadata for compact prefill.""" + + token_to_concept: Tensor + token_start_ids: Tensor + token_counts: Tensor + first_token_ids: Tensor + last_token_ids: Tensor + short_concept_mask: Tensor + + +def _flatten_decode_position_ids(position_ids: Tensor, batch_size: int, device: torch.device) -> Tensor: + """Normalize decode position ids to one absolute position per batch row.""" + if position_ids.dim() == 0: + position_ids = position_ids.view(1) + if position_ids.dim() == 1: + return position_ids.to(device=device, dtype=torch.long) + position_ids = position_ids.reshape(-1) + if position_ids.numel() == batch_size: + return position_ids.to(device=device, dtype=torch.long) + assert position_ids.numel() % batch_size == 0, ( + f'Cannot map position_ids with {position_ids.numel()} elements to batch size {batch_size}.') + return position_ids.reshape(-1, batch_size)[-1].to(device=device, dtype=torch.long) class DefaultConceptLMRuntimeOpsImpl(ConceptLMRuntimeOpsImpl): """Torch fallback implementation of ConceptLM runtime operations.""" + @staticmethod + def concept_count_from_seq_len(seq_len: int, chunk_size: int) -> int: + """Return reference ConceptLM chunk count for one request length.""" + seq_len = int(seq_len or 0) + if seq_len <= 0: + return 0 + if seq_len < chunk_size: + return 1 + return seq_len // chunk_size + + @staticmethod + def concept_counts_from_q_seqlens(q_seqlens: Tensor, chunk_size: int) -> Tensor: + """Vectorized version of ``concept_count_from_seq_len``.""" + counts = torch.div(q_seqlens, chunk_size, rounding_mode='floor').clamp(min=1) + return torch.where(q_seqlens > 0, counts, torch.zeros_like(q_seqlens)) + + @staticmethod + def repeat_slot_ids(token_pos: Tensor, chunk_size: int, shift_feature: bool) -> Tensor: + """Return local concept slot read by each token after shift + semantics.""" + if shift_feature: + return torch.div(token_pos + 1, chunk_size, rounding_mode='floor') - 1 + return torch.div(token_pos, chunk_size, rounding_mode='floor') - 1 + + def flatten_decode_position_ids(self, position_ids: Tensor, batch_size: int, device: torch.device) -> Tensor: + """Normalize decode position ids to one absolute position per batch + row.""" + return _flatten_decode_position_ids(position_ids, batch_size, device) + + def build_decode_metadata(self, + position_ids: Tensor, + state_ids: Tensor | None, + batch_size: int, + device: torch.device) -> ConceptDecodeMetadata: + """Build fixed-shape decode metadata from engine state ids.""" + if state_ids is None: + raise RuntimeError('ConceptLM decode requires state_ids.') + state_ids = state_ids.to(device=device, dtype=torch.long).reshape(-1) + if state_ids.numel() != batch_size: + raise ValueError(f'Expected {batch_size} decode state ids, got {state_ids.numel()}.') + valid_state_mask = state_ids >= 0 + return ConceptDecodeMetadata( + position_ids=position_ids, + state_ids=state_ids, + safe_state_ids=state_ids.clamp(min=0), + valid_state_mask=valid_state_mask, + ) + + @staticmethod + def select_decode_state_rows(state_cache: Tensor, decode_metadata: ConceptDecodeMetadata) -> Tensor: + """Gather state-cache rows and zero out padded decode rows.""" + rows = state_cache.index_select(0, decode_metadata.safe_state_ids) + mask_shape = (decode_metadata.valid_state_mask.size(0), ) + (1, ) * (rows.dim() - 1) + valid_mask = decode_metadata.valid_state_mask.view(mask_shape) + return torch.where(valid_mask, rows, torch.zeros_like(rows)) + + def select_decode_last_state_rows(self, + last_state: Tensor | None, + last_final_state: Tensor | None, + last_raw_states: Tensor | None, + decode_metadata: ConceptDecodeMetadata) -> tuple[Tensor, Tensor]: + """Gather packed last-concept state rows once and return final/raw + views.""" + if last_state is not None: + rows = self.select_decode_state_rows(last_state, decode_metadata) + return rows[:, 0], rows[:, 1:] + if last_final_state is None or last_raw_states is None: + raise RuntimeError('ConceptLM decode requires cached last concept states.') + + return ( + self.select_decode_state_rows(last_final_state, decode_metadata), + self.select_decode_state_rows(last_raw_states, decode_metadata), + ) + + def decode_concept_read_mask(self, decode_metadata: ConceptDecodeMetadata) -> Tensor: + """Return rows whose current decode token should read a cached + concept.""" + repeat_slots = self.repeat_slot_ids( + decode_metadata.position_ids, + self.chunk_size, + self.shift_feature, + ) + return decode_metadata.valid_state_mask & (repeat_slots >= 0) + + def build_concept_decode_metadata_static(self, token_attn_metadata: Any, + decode_metadata: ConceptDecodeMetadata): + """Build fixed-shape concept-stream decode metadata.""" + device = decode_metadata.position_ids.device + batch_size = decode_metadata.position_ids.numel() + q_seqlens = token_attn_metadata.q_seqlens + q_start_loc = token_attn_metadata.q_start_loc + kv_seqlens = token_attn_metadata.kv_seqlens + q_dtype = q_seqlens.dtype + q_start_dtype = q_start_loc.dtype + kv_dtype = kv_seqlens.dtype + + concept_q_seqlens = torch.ones((batch_size, ), dtype=q_dtype, device=device) + concept_q_start_loc = torch.arange(batch_size, dtype=q_start_dtype, device=device) + concept_cu_seqlens = F.pad(torch.cumsum(concept_q_seqlens, dim=0, dtype=torch.int32), (1, 0)) + concept_kv_seqlens = torch.div( + decode_metadata.position_ids + 1, + self.chunk_size, + rounding_mode='floor', + ).clamp(min=1).to(dtype=kv_dtype) + + updates = dict( + is_decoding=True, + block_offsets=token_attn_metadata.block_offsets, + q_start_loc=concept_q_start_loc, + q_seqlens=concept_q_seqlens, + kv_seqlens=concept_kv_seqlens, + cu_seqlens_q=concept_cu_seqlens, + cu_seqlens_k=concept_cu_seqlens, + ) + if hasattr(token_attn_metadata, 'kv_start_loc'): + updates['kv_start_loc'] = concept_kv_seqlens - concept_q_seqlens.to(dtype=concept_kv_seqlens.dtype) + if hasattr(token_attn_metadata, 'kv_flatten_size'): + updates['kv_flatten_size'] = batch_size + if hasattr(token_attn_metadata, 'max_q_seqlen'): + updates['max_q_seqlen'] = 1 + if hasattr(token_attn_metadata, 'max_kv_seqlen'): + updates['max_kv_seqlen'] = getattr(token_attn_metadata, 'max_kv_seqlen') + if hasattr(token_attn_metadata, 'scheduler_metadata'): + updates['scheduler_metadata'] = None + if hasattr(token_attn_metadata, 'tile_scheduler_metadata'): + updates['tile_scheduler_metadata'] = None + if hasattr(token_attn_metadata, 'num_splits'): + updates['num_splits'] = None + if hasattr(token_attn_metadata, 'fill_seqlens'): + updates['fill_seqlens'] = concept_q_seqlens + return replace(token_attn_metadata, **updates) + + def build_concept_chunk_input( + self, + source_states: Tensor, + token_attn_metadata: Any, + position_ids: Tensor, + state_ids: Tensor | None = None, + chunk_source_state_cache: Tensor | None = None, + ) -> ConceptChunkInput: + """Build concept-predictor source rows for prefill or decode. + + The caller provides source states in the same layout for both phases: + ``[token_or_batch, num_sources, hidden]``. Runtime metadata chooses the + compact prefill merge or fixed-shape decode accumulator update. + """ + if getattr(token_attn_metadata, 'is_decoding', False): + if chunk_source_state_cache is None: + raise RuntimeError('ConceptLM decode requires chunk source state cache.') + batch_size = source_states.size(0) + position_ids = self.flatten_decode_position_ids(position_ids, batch_size, source_states.device) + decode_metadata = self.build_decode_metadata( + position_ids, + state_ids, + batch_size, + source_states.device, + ) + concept_states, update_mask = self.decode_chunk_state_update( + chunk_source_state_cache, + source_states, + decode_metadata.state_ids, + decode_metadata.position_ids, + self.chunk_size, + self.merge_method, + ) + concept_attn_metadata = self.build_concept_decode_metadata_static(token_attn_metadata, decode_metadata) + concept_position_ids = self.decode_concept_position_ids(decode_metadata.position_ids) + return ConceptChunkInput( + source_states=concept_states, + position_ids=concept_position_ids, + attn_metadata=concept_attn_metadata, + state_ids=decode_metadata.state_ids, + update_mask=update_mask, + decode_metadata=decode_metadata, + ) + + prefill_metadata = self.build_prefill_metadata(token_attn_metadata, position_ids) + concept_states = self.prefill_chunk_state_update(source_states, prefill_metadata) + concept_attn_metadata = self.build_concept_prefill_metadata(token_attn_metadata, prefill_metadata) + return ConceptChunkInput( + source_states=concept_states, + position_ids=prefill_metadata.concept_position_ids, + attn_metadata=concept_attn_metadata, + state_ids=state_ids, + prefill_metadata=prefill_metadata, + ) + + def decode_concept_position_ids(self, position_ids: Tensor) -> Tensor: + """Return reference RoPE positions for emitted decode concept rows.""" + return (position_ids - self.chunk_size + 1).clamp(min=0) + + def _get_max_concepts_per_request(self, token_attn_metadata: Any, concept_q_seqlens: Tensor) -> int: + """Return per-request concept attention bound without hidden context + access.""" + max_q_seqlen = getattr(token_attn_metadata, 'max_q_seqlen', None) + if max_q_seqlen is not None: + return self.concept_count_from_seq_len(int(max_q_seqlen), self.chunk_size) + + # Test/direct-call fallback. Serving should use the scheduler-provided + # Python max_q_seqlen above, as in DSV4 metadata construction. + return int(concept_q_seqlens.max().item()) if concept_q_seqlens.numel() > 0 else 0 + + @staticmethod + def _build_prefill_token_layout(token_attn_metadata: Any, position_ids: Tensor) -> _PrefillTokenLayout: + """Build packed token-stream positions from prefill attention + metadata.""" + q_seqlens = token_attn_metadata.q_seqlens + q_start_loc = getattr(token_attn_metadata, 'q_start_loc', None) + if q_start_loc is None: + q_start_loc = F.pad(torch.cumsum(q_seqlens, dim=0, dtype=torch.int32), (1, 0))[:-1] + + total_tokens = int(position_ids.numel()) + + q_seqlens_long = q_seqlens.to(dtype=torch.long, device=position_ids.device) + q_start_loc_long = q_start_loc.to(dtype=torch.long, device=position_ids.device) + cu_q_seqlens = getattr(token_attn_metadata, 'cu_seqlens_q', None) + if cu_q_seqlens is None: + cu_q_seqlens = F.pad(torch.cumsum(q_seqlens, dim=0, dtype=torch.int32), (1, 0)) + cu_q_seqlens_long = cu_q_seqlens.to(dtype=torch.long, device=position_ids.device) + + token_ids = torch.arange(total_tokens, dtype=torch.long, device=position_ids.device) + token_seq = torch.searchsorted(cu_q_seqlens_long[1:], token_ids, right=True) + token_pos = token_ids - cu_q_seqlens_long[token_seq] + + return _PrefillTokenLayout( + q_seqlens=q_seqlens, + q_start_loc=q_start_loc, + q_seqlens_long=q_seqlens_long, + q_start_loc_long=q_start_loc_long, + token_seq=token_seq, + token_pos=token_pos, + total_tokens=total_tokens, + ) + + def _build_prefill_concept_layout(self, token_attn_metadata: Any, position_ids: Tensor, + token_layout: _PrefillTokenLayout) -> _PrefillConceptLayout: + """Build compact chunk-token positions used by concept attention.""" + concept_q_seqlens_long = self.concept_counts_from_q_seqlens(token_layout.q_seqlens_long, self.chunk_size) + concept_q_seqlens = concept_q_seqlens_long.to(dtype=token_layout.q_seqlens.dtype, + device=token_layout.q_seqlens.device) + concept_q_start_loc = F.pad(torch.cumsum(concept_q_seqlens, dim=0, dtype=torch.int32), (1, 0))[:-1] + concept_q_start_loc_long = concept_q_start_loc.to(dtype=torch.long, device=position_ids.device) + concept_cu_seqlens_long = F.pad(torch.cumsum(concept_q_seqlens_long, dim=0), (1, 0)) + + # TODO: remove this eager compact-size scalar read by moving ConceptLM + # concept-stream allocation/compaction into the engine/backend contract, + # like DSV4's precomputed metadata and Qwen3.5's state-cache metadata. + num_concepts_total = int(concept_cu_seqlens_long[-1].item()) + concept_ids = torch.arange(num_concepts_total, dtype=torch.long, device=position_ids.device) + concept_seq = torch.searchsorted(concept_cu_seqlens_long[1:], concept_ids, right=True) + local_concept_ids = concept_ids - concept_cu_seqlens_long[concept_seq] + concept_token_start = token_layout.q_start_loc_long[concept_seq] + local_concept_ids * self.chunk_size + concept_position_ids = position_ids[concept_token_start] + max_concepts_per_request = self._get_max_concepts_per_request(token_attn_metadata, concept_q_seqlens_long) + + return _PrefillConceptLayout( + q_seqlens=concept_q_seqlens, + q_seqlens_long=concept_q_seqlens_long, + q_start_loc=concept_q_start_loc, + q_start_loc_long=concept_q_start_loc_long, + seq=concept_seq, + local_ids=local_concept_ids, + position_ids=concept_position_ids, + num_total=num_concepts_total, + max_per_request=max_concepts_per_request, + ) + + def _build_prefill_repeat_ids(self, token_layout: _PrefillTokenLayout, + concept_layout: _PrefillConceptLayout) -> Tensor: + """Map packed token rows to the concept row visible after shift.""" + seq_concept_start = concept_layout.q_start_loc_long[token_layout.token_seq] + seq_concept_count = concept_layout.q_seqlens_long[token_layout.token_seq] + repeat_slots = self.repeat_slot_ids(token_layout.token_pos, self.chunk_size, self.shift_feature) + valid_repeat = (repeat_slots >= 0) & (repeat_slots < seq_concept_count) + return torch.where( + valid_repeat, + seq_concept_start + repeat_slots, + torch.full_like(repeat_slots, -1), + ) + + def _build_prefill_merge_layout(self, token_layout: _PrefillTokenLayout, + concept_layout: _PrefillConceptLayout) -> _PrefillMergeLayout: + """Build token-to-concept merge metadata for mean/first/last modes.""" + seq_concept_start = concept_layout.q_start_loc_long[token_layout.token_seq] + seq_concept_count = concept_layout.q_seqlens_long[token_layout.token_seq] + merge_slots = torch.div(token_layout.token_pos, self.chunk_size, rounding_mode='floor') + valid_merge = (merge_slots >= 0) & (merge_slots < seq_concept_count) + merge_token_to_concept = torch.where( + valid_merge, + seq_concept_start + merge_slots, + torch.full_like(merge_slots, -1), + ) + + concept_seq_len = token_layout.q_seqlens_long[concept_layout.seq] + merge_start_pos = concept_layout.local_ids * self.chunk_size + merge_counts_long = (concept_seq_len - merge_start_pos).clamp(min=0) + merge_counts_long = torch.minimum( + merge_counts_long, + torch.full_like(merge_counts_long, self.chunk_size), + ) + merge_token_counts = merge_counts_long.to(dtype=torch.int32) + merge_token_start_ids = token_layout.q_start_loc_long[concept_layout.seq] + merge_start_pos + merge_first_token_ids = merge_token_start_ids + merge_last_token_ids = merge_token_start_ids + merge_counts_long.clamp(min=1) - 1 + merge_short_concept_mask = merge_counts_long < self.chunk_size + + return _PrefillMergeLayout( + token_to_concept=merge_token_to_concept, + token_start_ids=merge_token_start_ids, + token_counts=merge_token_counts, + first_token_ids=merge_first_token_ids, + last_token_ids=merge_last_token_ids, + short_concept_mask=merge_short_concept_mask, + ) + + def build_prefill_metadata(self, token_attn_metadata: Any, position_ids: Tensor) -> ConceptPrefillMetadata: + """Build packed token-to-concept metadata for batched prefill.""" + token_layout = self._build_prefill_token_layout(token_attn_metadata, position_ids) + concept_layout = self._build_prefill_concept_layout(token_attn_metadata, position_ids, token_layout) + token_to_concept = self._build_prefill_repeat_ids(token_layout, concept_layout) + merge_layout = self._build_prefill_merge_layout(token_layout, concept_layout) + + return ConceptPrefillMetadata( + token_q_seqlens=token_layout.q_seqlens, + token_q_start_loc=token_layout.q_start_loc, + concept_q_seqlens=concept_layout.q_seqlens, + concept_q_start_loc=concept_layout.q_start_loc, + concept_position_ids=concept_layout.position_ids, + merge_token_to_concept=merge_layout.token_to_concept, + merge_token_start_ids=merge_layout.token_start_ids, + merge_token_counts=merge_layout.token_counts, + merge_first_token_ids=merge_layout.first_token_ids, + merge_last_token_ids=merge_layout.last_token_ids, + merge_short_concept_mask=merge_layout.short_concept_mask, + token_to_concept=token_to_concept, + num_tokens_total=token_layout.total_tokens, + num_concepts_total=concept_layout.num_total, + max_concepts_per_request=concept_layout.max_per_request, + ) + + @staticmethod + def _merge_chunks_mean_packed(hidden_states: Tensor, prefill_metadata: ConceptPrefillMetadata) -> Tensor: + """Mean-pool packed token states by precomputed concept ids.""" + merge_token_to_concept = prefill_metadata.merge_token_to_concept.to(device=hidden_states.device) + valid_merge = (merge_token_to_concept >= 0).to(dtype=hidden_states.dtype).unsqueeze(-1) + safe_merge_ids = merge_token_to_concept.clamp(min=0) + merged = hidden_states.new_zeros((prefill_metadata.num_concepts_total, hidden_states.size(-1))) + merged.index_add_(0, safe_merge_ids, hidden_states * valid_merge) + counts = prefill_metadata.merge_token_counts.clamp(min=1).to(device=hidden_states.device, + dtype=hidden_states.dtype) + return merged / counts.unsqueeze(-1) + + def merge_chunks_packed(self, hidden_states: Tensor, prefill_metadata: ConceptPrefillMetadata) -> Tensor: + """Merge packed token states into packed per-request concept states.""" + if self.merge_method == 'first': + merged = hidden_states[prefill_metadata.merge_first_token_ids] + short_mean = self._merge_chunks_mean_packed(hidden_states, prefill_metadata) + return torch.where(prefill_metadata.merge_short_concept_mask[:, None], short_mean, merged) + if self.merge_method == 'last': + merged = hidden_states[prefill_metadata.merge_last_token_ids] + short_mean = self._merge_chunks_mean_packed(hidden_states, prefill_metadata) + return torch.where(prefill_metadata.merge_short_concept_mask[:, None], short_mean, merged) + + return self._merge_chunks_mean_packed(hidden_states, prefill_metadata) + + def prefill_chunk_state_update(self, source_states: Tensor, prefill_metadata: ConceptPrefillMetadata) -> Tensor: + """Merge prefill source states to the unified concept-source layout.""" + if source_states.dim() == 2: + return self.merge_chunks_packed(source_states, prefill_metadata) + if source_states.dim() != 3: + raise ValueError(f'ConceptLM prefill source states must be 2-D or 3-D, got {tuple(source_states.shape)}.') + chunks = [self.merge_chunks_packed(source_states[:, source_id], prefill_metadata) + for source_id in range(source_states.size(1))] + return torch.stack(tuple(chunks), dim=1) + + @staticmethod + def _gather_zero_prefixed_concepts(concept_states_with_zero: Tensor, + prefill_metadata: ConceptPrefillMetadata) -> Tensor: + """Gather zero-prefixed concept rows to packed token rows.""" + token_to_concept = prefill_metadata.token_to_concept.to(device=concept_states_with_zero.device) + gather_ids = torch.clamp(token_to_concept + 1, min=0) + return concept_states_with_zero[gather_ids] + + def repeat_shift_packed(self, concept_states: Tensor, prefill_metadata: ConceptPrefillMetadata) -> Tensor: + """Gather packed concept states back to packed token states.""" + concept_states_with_zero = torch.cat((torch.zeros_like(concept_states[:1]), concept_states), dim=0) + return self._gather_zero_prefixed_concepts(concept_states_with_zero, prefill_metadata) + + def repeat_shift_source_states_packed(self, concept_states_with_zero: Tensor, + prefill_metadata: ConceptPrefillMetadata) -> Tensor: + """Gather zero-prefixed packed concept source states to token rows.""" + return self._gather_zero_prefixed_concepts(concept_states_with_zero, prefill_metadata) + + def begin_concept_forward(self, chunk_input: ConceptChunkInput, + runtime_caches: ConceptRuntimeCaches) -> ConceptForwardContext: + """Prepare transient state before running the concept predictor.""" + if not chunk_input.is_decoding: + return ConceptForwardContext() + + concept_past_key_values = runtime_caches.concept_past_key_values + if concept_past_key_values is None: + raise RuntimeError('ConceptLM decode requires concept KV caches.') + decode_metadata = chunk_input.decode_metadata + if decode_metadata is None: + raise RuntimeError('ConceptLM decode input is missing decode metadata.') + saved_kv = self.snapshot_decode_concept_kv(concept_past_key_values, chunk_input.attn_metadata) + + previous_final_state = None + previous_raw_states = None + if not self.shift_feature: + previous_final_state, previous_raw_states = self.select_decode_last_state_rows( + runtime_caches.last_state, + runtime_caches.last_final_state, + runtime_caches.last_raw_states, + decode_metadata, + ) + + return ConceptForwardContext( + saved_kv=saved_kv, + previous_final_state=previous_final_state, + previous_raw_states=previous_raw_states, + ) + + def end_concept_forward( + self, + chunk_input: ConceptChunkInput, + runtime_caches: ConceptRuntimeCaches, + forward_context: ConceptForwardContext, + source_states: Tensor, + predicted_vectors: Tensor, + concept_raw_states: list[Tensor], + ) -> None: + """Commit concept-predictor side effects for prefill or decode.""" + if not chunk_input.is_decoding: + prefill_metadata = chunk_input.prefill_metadata + if prefill_metadata is None: + raise RuntimeError('ConceptLM prefill input is missing prefill metadata.') + self.write_prefill_state_caches( + runtime_caches.chunk_source_state, + runtime_caches.last_raw_states, + runtime_caches.last_final_state, + chunk_input.state_ids, + prefill_metadata, + source_states, + predicted_vectors, + concept_raw_states, + ) + return + + concept_past_key_values = runtime_caches.concept_past_key_values + if concept_past_key_values is None: + raise RuntimeError('ConceptLM decode requires concept KV caches.') + if forward_context.saved_kv is None: + raise RuntimeError('ConceptLM decode forward context is missing KV snapshot.') + if chunk_input.decode_metadata is None or chunk_input.update_mask is None: + raise RuntimeError('ConceptLM decode input is missing metadata or update mask.') + if runtime_caches.last_raw_states is None or runtime_caches.last_final_state is None: + raise RuntimeError('ConceptLM decode requires cached last concept states.') + self.restore_decode_concept_kv( + concept_past_key_values, + chunk_input.attn_metadata, + forward_context.saved_kv, + ~chunk_input.update_mask, + ) + self.write_decode_concept_states( + runtime_caches.last_raw_states, + runtime_caches.last_final_state, + predicted_vectors, + concept_raw_states, + chunk_input.decode_metadata.state_ids, + chunk_input.update_mask, + ) + + def build_decoder_concept_input( + self, + chunk_input: ConceptChunkInput, + runtime_caches: ConceptRuntimeCaches, + forward_context: ConceptForwardContext, + predicted_vectors: Tensor, + concept_raw_states: list[Tensor], + ) -> ConceptDecoderInput: + """Build token-decoder concept inputs from committed concept state.""" + if not chunk_input.is_decoding: + prefill_metadata = chunk_input.prefill_metadata + if prefill_metadata is None: + raise RuntimeError('ConceptLM prefill input is missing prefill metadata.') + final_state = self.repeat_shift_packed(predicted_vectors, prefill_metadata) + raw_states = self.stack_concept_raw_states(concept_raw_states) + zero_chunk = torch.zeros_like(raw_states[:1]) + route_states = self.repeat_shift_source_states_packed( + torch.cat((zero_chunk, raw_states), dim=0), + prefill_metadata, + ) + return ConceptDecoderInput(final_state=final_state, route_states=route_states) + + decode_metadata = chunk_input.decode_metadata + if decode_metadata is None: + raise RuntimeError('ConceptLM decode input is missing decode metadata.') + if self.shift_feature: + final_state, route_states = self.select_decode_last_state_rows( + runtime_caches.last_state, + runtime_caches.last_final_state, + runtime_caches.last_raw_states, + decode_metadata, + ) + else: + final_state = forward_context.previous_final_state + route_states = forward_context.previous_raw_states + if final_state is None or route_states is None: + raise RuntimeError('ConceptLM decode forward context is missing previous concept states.') + + concept_read_mask = self.decode_concept_read_mask(decode_metadata) + final_state = torch.where(concept_read_mask.view(-1, 1), final_state, torch.zeros_like(final_state)) + route_states = torch.where(concept_read_mask.view(-1, 1, 1), route_states, torch.zeros_like(route_states)) + return ConceptDecoderInput(final_state=final_state, route_states=route_states) + + def build_concept_prefill_metadata(self, token_attn_metadata: Any, prefill_metadata: ConceptPrefillMetadata): + """Build chunk-stream attention metadata for packed prefill.""" + concept_q_seqlens = prefill_metadata.concept_q_seqlens + concept_q_start_loc = prefill_metadata.concept_q_start_loc + concept_cu_seqlens = F.pad( + torch.cumsum(concept_q_seqlens, dim=0, dtype=torch.int32), + (1, 0), + ) + max_concept_seqlen = int(prefill_metadata.max_concepts_per_request) + + updates = dict( + is_decoding=False, + q_start_loc=concept_q_start_loc, + q_seqlens=concept_q_seqlens, + kv_seqlens=concept_q_seqlens, + cu_seqlens_q=concept_cu_seqlens, + cu_seqlens_k=concept_cu_seqlens, + ) + if hasattr(token_attn_metadata, 'kv_start_loc'): + updates['kv_start_loc'] = concept_q_start_loc + if hasattr(token_attn_metadata, 'kv_flatten_size'): + updates['kv_flatten_size'] = int(prefill_metadata.num_concepts_total) + if hasattr(token_attn_metadata, 'max_q_seqlen'): + updates['max_q_seqlen'] = max_concept_seqlen + if hasattr(token_attn_metadata, 'max_kv_seqlen'): + updates['max_kv_seqlen'] = max_concept_seqlen + return replace(token_attn_metadata, **updates) + + def merge_prefill_tail_chunk_states(self, source_states: Tensor, + prefill_metadata: ConceptPrefillMetadata) -> Tensor: + """Build per-request partial chunk accumulator rows after prefill.""" + device = source_states.device + q_seqlens = prefill_metadata.token_q_seqlens.to(device=device, dtype=torch.long) + q_start_loc = prefill_metadata.token_q_start_loc.to(device=device, dtype=torch.long) + batch_size = q_seqlens.size(0) + tail_lens = torch.remainder(q_seqlens, self.chunk_size) + tail_lens = torch.where(q_seqlens < self.chunk_size, q_seqlens, tail_lens) + tail_lens = torch.where(q_seqlens > 0, tail_lens, torch.zeros_like(tail_lens)) + has_tail = tail_lens > 0 + + tail_rows = source_states.new_zeros((batch_size, source_states.size(1), source_states.size(2))) + if source_states.size(0) == 0: + return tail_rows + if self.merge_method == 'first': + token_ids = q_start_loc + q_seqlens - tail_lens + token_ids = token_ids.clamp(min=0, max=max(source_states.size(0) - 1, 0)) + rows = source_states[token_ids] + return torch.where(has_tail.view(batch_size, 1, 1), rows, tail_rows) + if self.merge_method == 'last': + token_ids = q_start_loc + q_seqlens - 1 + token_ids = token_ids.clamp(min=0, max=max(source_states.size(0) - 1, 0)) + rows = source_states[token_ids] + return torch.where(has_tail.view(batch_size, 1, 1), rows, tail_rows) + + token_ids = torch.arange(prefill_metadata.num_tokens_total, dtype=torch.long, device=device) + cu_q_seqlens = F.pad(torch.cumsum(q_seqlens, dim=0), (1, 0)) + token_seq = torch.searchsorted(cu_q_seqlens[1:], token_ids, right=True) + token_pos = token_ids - cu_q_seqlens[token_seq] + token_tail_start = q_seqlens[token_seq] - tail_lens[token_seq] + valid_tail = (tail_lens[token_seq] > 0) & (token_pos >= token_tail_start) + weighted_source = source_states * valid_tail.to(dtype=source_states.dtype).view(-1, 1, 1) + tail_rows.index_add_(0, token_seq, weighted_source) + return tail_rows + + @staticmethod + def stack_concept_raw_states(concept_raw_states: list[Tensor]) -> Tensor: + """Stack raw concept-layer states to ``[rows, concept_layers, + hidden]``.""" + return torch.stack(tuple(concept_raw_states), dim=1) + + def write_prefill_state_caches( + self, + chunk_source_state: Tensor | None, + last_raw_states: Tensor | None, + last_final_state: Tensor | None, + state_ids: Tensor | None, + prefill_metadata: ConceptPrefillMetadata, + source_states: Tensor, + predicted_vectors: Tensor, + concept_raw_states: list[Tensor], + ) -> None: + """Seed decode state caches from a completed prefill forward.""" + return self.write_prefill_state_caches_eager( + chunk_source_state, + last_raw_states, + last_final_state, + state_ids, + prefill_metadata, + source_states, + predicted_vectors, + concept_raw_states, + ) + + def write_prefill_state_caches_eager( + self, + chunk_source_state: Tensor | None, + last_raw_states: Tensor | None, + last_final_state: Tensor | None, + state_ids: Tensor | None, + prefill_metadata: ConceptPrefillMetadata, + source_states: Tensor, + predicted_vectors: Tensor, + concept_raw_states: list[Tensor], + ) -> None: + """Seed decode state caches from a completed prefill forward.""" + if state_ids is None: + return + if chunk_source_state is None or last_raw_states is None or last_final_state is None: + return + + state_ids = state_ids.to(device=source_states.device, dtype=torch.long).reshape(-1) + valid_indices = torch.nonzero(state_ids >= 0, as_tuple=False).flatten() + if valid_indices.numel() == 0: + return + + valid_state_ids = state_ids.index_select(0, valid_indices) + tail_rows = self.merge_prefill_tail_chunk_states(source_states, prefill_metadata) + chunk_source_state.index_copy_(0, valid_state_ids, + tail_rows.index_select(0, valid_indices).to(dtype=chunk_source_state.dtype)) + + concept_counts = prefill_metadata.concept_q_seqlens.to(device=source_states.device, dtype=torch.long) + concept_start = prefill_metadata.concept_q_start_loc.to(device=source_states.device, dtype=torch.long) + concept_valid_mask = (state_ids >= 0) & (concept_counts > 0) + concept_indices = torch.nonzero(concept_valid_mask, as_tuple=False).flatten() + if concept_indices.numel() == 0: + return + + concept_state_ids = state_ids.index_select(0, concept_indices) + last_concept_ids = concept_start.index_select(0, concept_indices) + concept_counts.index_select( + 0, concept_indices) - 1 + last_final_rows = predicted_vectors.index_select(0, last_concept_ids) + last_final_state.index_copy_(0, concept_state_ids, last_final_rows.to(dtype=last_final_state.dtype)) + raw_rows = self.stack_concept_raw_states(concept_raw_states).index_select(0, last_concept_ids) + last_raw_states.index_copy_(0, concept_state_ids, raw_rows.to(dtype=last_raw_states.dtype)) + + def snapshot_decode_concept_kv(self, concept_past_key_values: list[list[Tensor]], + concept_attn_metadata: Any) -> list[tuple[Tensor, Tensor]]: + """Snapshot concept KV slots that dummy non-boundary rows may + overwrite.""" + return [ + self.decode_kv_cache_snapshot( + k_cache, + v_cache, + concept_attn_metadata.block_offsets, + concept_attn_metadata.kv_seqlens, + ) + for k_cache, v_cache in concept_past_key_values + ] + + def restore_decode_concept_kv(self, concept_past_key_values: list[list[Tensor]], concept_attn_metadata: Any, + saved_kv: list[tuple[Tensor, Tensor]], restore_mask: Tensor) -> None: + """Restore concept KV slots for non-boundary and padded rows.""" + for (k_cache, v_cache), (saved_k, saved_v) in zip(concept_past_key_values, saved_kv): + self.decode_kv_cache_restore( + k_cache, + v_cache, + saved_k, + saved_v, + concept_attn_metadata.block_offsets, + concept_attn_metadata.kv_seqlens, + restore_mask, + ) + + def write_decode_concept_states( + self, + last_raw_state_cache: Tensor, + last_final_state_cache: Tensor, + predicted_vectors: Tensor, + raw_states: list[Tensor], + state_ids: Tensor, + update_mask: Tensor, + ) -> None: + """Write newly emitted concept states to persistent decode caches.""" + raw_rows = self.stack_concept_raw_states(raw_states) + self.decode_concept_state_update( + last_raw_state_cache, + last_final_state_cache, + predicted_vectors, + raw_rows, + state_ids, + update_mask, + ) + @staticmethod def _decode_kv_cache_rows(k_cache: Tensor, block_offsets: Tensor, diff --git a/lmdeploy/pytorch/kernels/cuda/conceptlm.py b/lmdeploy/pytorch/kernels/cuda/conceptlm.py index b329b26a1d..0c3b959734 100644 --- a/lmdeploy/pytorch/kernels/cuda/conceptlm.py +++ b/lmdeploy/pytorch/kernels/cuda/conceptlm.py @@ -6,6 +6,164 @@ import triton.language as tl +@triton.jit +def _prefill_chunk_state_update_kernel( + source_states, + concept_states, + token_start_ids, + token_counts, + source_stride_t, + source_stride_s, + source_stride_h, + out_stride_c, + out_stride_s, + out_stride_h, + HIDDEN: tl.constexpr, + TOTAL_ELEMS: tl.constexpr, + CHUNK_SIZE: tl.constexpr, + MERGE_METHOD: tl.constexpr, + BLOCK: tl.constexpr, +): + """Merge contiguous token chunks into compact concept-source rows.""" + tile_id = tl.program_id(0) + concept_id = tl.program_id(1) + offs = tile_id * BLOCK + tl.arange(0, BLOCK) + valid_elem = offs < TOTAL_ELEMS + + source_id = offs // HIDDEN + hidden_id = offs - source_id * HIDDEN + + token_start = tl.load(token_start_ids + concept_id) + token_count = tl.load(token_counts + concept_id) + has_token = token_count > 0 + + acc = tl.zeros((BLOCK, ), dtype=tl.float32) + for token_offset in range(CHUNK_SIZE): + load_mask = valid_elem & (token_offset < token_count) + ptrs = (source_states + (token_start + token_offset) * source_stride_t + source_id * source_stride_s + + hidden_id * source_stride_h) + values = tl.load(ptrs, mask=load_mask, other=0.0).to(tl.float32) + acc += values + + denom = tl.maximum(token_count, 1).to(tl.float32) + mean_values = acc / denom + if MERGE_METHOD == 1: # first; short prompts keep reference mean-pooling + first_ptrs = (source_states + token_start * source_stride_t + source_id * source_stride_s + + hidden_id * source_stride_h) + first_values = tl.load(first_ptrs, mask=valid_elem & has_token, other=0.0).to(tl.float32) + out_values = tl.where(token_count < CHUNK_SIZE, mean_values, first_values) + elif MERGE_METHOD == 2: # last; short prompts keep reference mean-pooling + last_token = token_start + tl.maximum(token_count, 1) - 1 + last_ptrs = (source_states + last_token * source_stride_t + source_id * source_stride_s + + hidden_id * source_stride_h) + last_values = tl.load(last_ptrs, mask=valid_elem & has_token, other=0.0).to(tl.float32) + out_values = tl.where(token_count < CHUNK_SIZE, mean_values, last_values) + else: + out_values = mean_values + + out_ptrs = concept_states + concept_id * out_stride_c + source_id * out_stride_s + hidden_id * out_stride_h + tl.store(out_ptrs, out_values, mask=valid_elem) + + +@triton.jit +def _prefill_state_cache_update_kernel( + chunk_state_cache, + last_raw_state_cache, + last_final_state_cache, + source_states, + predicted_vectors, + raw_states, + state_ids, + token_q_start_loc, + token_q_seqlens, + concept_q_start_loc, + concept_q_seqlens, + chunk_state_stride_n, + chunk_state_stride_s, + chunk_state_stride_h, + raw_cache_stride_n, + raw_cache_stride_l, + raw_cache_stride_h, + final_cache_stride_n, + final_cache_stride_h, + source_stride_t, + source_stride_s, + source_stride_h, + pred_stride_c, + pred_stride_h, + raw_stride_c, + raw_stride_l, + raw_stride_h, + HIDDEN: tl.constexpr, + SOURCE_ELEMS: tl.constexpr, + RAW_ELEMS: tl.constexpr, + CHUNK_SIZE: tl.constexpr, + MERGE_METHOD: tl.constexpr, + BLOCK: tl.constexpr, +): + """Seed decode state caches directly from prefill rows.""" + tile_id = tl.program_id(0) + batch_id = tl.program_id(1) + state_id = tl.load(state_ids + batch_id) + if state_id < 0: + return + + offs = tile_id * BLOCK + tl.arange(0, BLOCK) + q_start = tl.load(token_q_start_loc + batch_id).to(tl.int64) + q_len = tl.load(token_q_seqlens + batch_id).to(tl.int64) + tail_len = q_len % CHUNK_SIZE + tail_len = tl.where(q_len < CHUNK_SIZE, q_len, tail_len) + tail_len = tl.where(q_len > 0, tail_len, 0) + has_tail = tail_len > 0 + + source_mask = offs < SOURCE_ELEMS + source_id = offs // HIDDEN + source_hidden_id = offs - source_id * HIDDEN + tail_start = q_start + q_len - tail_len + + if MERGE_METHOD == 1: # first + first_ptrs = (source_states + tail_start * source_stride_t + source_id * source_stride_s + + source_hidden_id * source_stride_h) + source_values = tl.load(first_ptrs, mask=source_mask & has_tail, other=0.0).to(tl.float32) + elif MERGE_METHOD == 2: # last + last_ptrs = (source_states + (q_start + q_len - 1) * source_stride_t + source_id * source_stride_s + + source_hidden_id * source_stride_h) + source_values = tl.load(last_ptrs, mask=source_mask & has_tail, other=0.0).to(tl.float32) + else: + source_values = tl.zeros((BLOCK, ), dtype=tl.float32) + for token_offset in range(CHUNK_SIZE): + load_mask = source_mask & (token_offset < tail_len) + ptrs = (source_states + (tail_start + token_offset) * source_stride_t + source_id * source_stride_s + + source_hidden_id * source_stride_h) + values = tl.load(ptrs, mask=load_mask, other=0.0).to(tl.float32) + source_values += values + + chunk_ptrs = (chunk_state_cache + state_id * chunk_state_stride_n + source_id * chunk_state_stride_s + + source_hidden_id * chunk_state_stride_h) + tl.store(chunk_ptrs, source_values, mask=source_mask) + + concept_count = tl.load(concept_q_seqlens + batch_id).to(tl.int64) + has_concept = concept_count > 0 + concept_start = tl.load(concept_q_start_loc + batch_id).to(tl.int64) + last_concept_id = concept_start + concept_count - 1 + + final_hidden_id = offs + final_mask = (offs < HIDDEN) & has_concept + pred_ptrs = predicted_vectors + last_concept_id * pred_stride_c + final_hidden_id * pred_stride_h + final_ptrs = last_final_state_cache + state_id * final_cache_stride_n + final_hidden_id * final_cache_stride_h + final_values = tl.load(pred_ptrs, mask=final_mask, other=0.0) + tl.store(final_ptrs, final_values, mask=final_mask) + + raw_mask = (offs < RAW_ELEMS) & has_concept + raw_layer_id = offs // HIDDEN + raw_hidden_id = offs - raw_layer_id * HIDDEN + raw_ptrs = raw_states + last_concept_id * raw_stride_c + raw_layer_id * raw_stride_l + raw_hidden_id * raw_stride_h + raw_cache_ptrs = (last_raw_state_cache + state_id * raw_cache_stride_n + raw_layer_id * raw_cache_stride_l + + raw_hidden_id * raw_cache_stride_h) + raw_values = tl.load(raw_ptrs, mask=raw_mask, other=0.0) + tl.store(raw_cache_ptrs, raw_values, mask=raw_mask) + + @triton.jit def _decode_chunk_state_update_kernel( state_cache, @@ -259,6 +417,122 @@ def _merge_method_id(merge_method: str) -> int: return 0 +def prefill_chunk_state_update( + source_states: torch.Tensor, + token_start_ids: torch.Tensor, + token_counts: torch.Tensor, + num_concepts_total: int, + chunk_size: int, + merge_method: str, + block: int = 1024, +): + """Merge prefill source states into compact concept-source rows. + + Args: + source_states: ``[num_tokens, num_sources, hidden]``. + token_start_ids: first token row for each compact concept row. + token_counts: number of tokens merged into each concept row. + num_concepts_total: compact concept row count. + chunk_size: ConceptLM chunk size. + merge_method: ``meanpooling``, ``first``, or ``last``. + block: Triton vector width over ``num_sources * hidden``. + """ + assert source_states.is_cuda, 'ConceptLM prefill merge requires CUDA source states.' + assert source_states.dim() == 3 + num_sources = source_states.size(1) + hidden = source_states.size(2) + concept_states = source_states.new_empty((num_concepts_total, num_sources, hidden)) + if num_concepts_total == 0: + return concept_states + + total_elems = num_sources * hidden + token_start_ids = token_start_ids.to(device=source_states.device, dtype=torch.long) + token_counts = token_counts.to(device=source_states.device, dtype=torch.int32) + grid = (triton.cdiv(total_elems, block), num_concepts_total) + _prefill_chunk_state_update_kernel[grid]( + source_states, + concept_states, + token_start_ids, + token_counts, + *source_states.stride(), + *concept_states.stride(), + HIDDEN=hidden, + TOTAL_ELEMS=total_elems, + CHUNK_SIZE=int(chunk_size), + MERGE_METHOD=_merge_method_id(merge_method), + BLOCK=block, + num_warps=8, + ) + return concept_states + + +def prefill_state_cache_update( + chunk_state_cache: torch.Tensor, + last_raw_state_cache: torch.Tensor, + last_final_state_cache: torch.Tensor, + source_states: torch.Tensor, + predicted_vectors: torch.Tensor, + raw_states: torch.Tensor, + state_ids: torch.Tensor, + token_q_start_loc: torch.Tensor, + token_q_seqlens: torch.Tensor, + concept_q_start_loc: torch.Tensor, + concept_q_seqlens: torch.Tensor, + chunk_size: int, + merge_method: str, + block: int = 1024, +) -> None: + """Seed ConceptLM decode state caches from prefill in one CUDA op.""" + assert chunk_state_cache.is_cuda, 'ConceptLM prefill state-cache update requires CUDA caches.' + assert last_raw_state_cache.is_cuda and last_final_state_cache.is_cuda + assert source_states.is_cuda and predicted_vectors.is_cuda and raw_states.is_cuda + assert source_states.dim() == 3 and predicted_vectors.dim() == 2 and raw_states.dim() == 3 + assert chunk_state_cache.shape[1:] == source_states.shape[1:] + assert last_final_state_cache.size(1) == predicted_vectors.size(1) + assert last_raw_state_cache.shape[1:] == raw_states.shape[1:] + + batch_size = token_q_seqlens.numel() + if batch_size == 0: + return + + hidden = source_states.size(2) + source_elems = source_states.size(1) * hidden + raw_elems = raw_states.size(1) * raw_states.size(2) + max_elems = max(source_elems, hidden, raw_elems) + state_ids = state_ids.to(device=source_states.device, dtype=torch.long).reshape(-1) + token_q_start_loc = token_q_start_loc.to(device=source_states.device) + token_q_seqlens = token_q_seqlens.to(device=source_states.device) + concept_q_start_loc = concept_q_start_loc.to(device=source_states.device) + concept_q_seqlens = concept_q_seqlens.to(device=source_states.device) + grid = (triton.cdiv(max_elems, block), batch_size) + _prefill_state_cache_update_kernel[grid]( + chunk_state_cache, + last_raw_state_cache, + last_final_state_cache, + source_states, + predicted_vectors, + raw_states, + state_ids, + token_q_start_loc, + token_q_seqlens, + concept_q_start_loc, + concept_q_seqlens, + *chunk_state_cache.stride(), + *last_raw_state_cache.stride(), + *last_final_state_cache.stride(), + *source_states.stride(), + *predicted_vectors.stride(), + *raw_states.stride(), + HIDDEN=hidden, + SOURCE_ELEMS=source_elems, + RAW_ELEMS=raw_elems, + CHUNK_SIZE=int(chunk_size), + MERGE_METHOD=_merge_method_id(merge_method), + BLOCK=block, + num_warps=8, + ) + + def decode_chunk_state_update( chunk_source_state_cache: torch.Tensor, current_source_states: torch.Tensor, diff --git a/lmdeploy/pytorch/models/intern_ncp/modeling.py b/lmdeploy/pytorch/models/intern_ncp/modeling.py index 702aec35be..3cd07b82ca 100644 --- a/lmdeploy/pytorch/models/intern_ncp/modeling.py +++ b/lmdeploy/pytorch/models/intern_ncp/modeling.py @@ -1,6 +1,6 @@ # Copyright (c) OpenMMLab. All rights reserved. from collections.abc import Iterable, Mapping -from dataclasses import dataclass, replace +from dataclasses import dataclass from typing import Any import torch @@ -8,9 +8,9 @@ from transformers.configuration_utils import PretrainedConfig from lmdeploy.pytorch.backends.conceptlm import ( - ConceptChunkStateUpdateResult, - ConceptDecodeMetadata, - ConceptPrefillMetadata, + ConceptChunkInput, + ConceptDecoderInput, + ConceptRuntimeCaches, ) from lmdeploy.pytorch.model_inputs import StepContext, StepContextManager from lmdeploy.pytorch.nn import ConceptLMRuntimeOps @@ -53,37 +53,6 @@ class _ConceptPredictorRequest: attn_metadata: Any -@dataclass -class _DecoderConceptInput: - """Concept inputs consumed by the decoder stack.""" - - final_state: torch.Tensor - raw_states: list[torch.Tensor] | None = None - prefill_metadata: ConceptPrefillMetadata | None = None - decode_states: torch.Tensor | None = None - - @classmethod - def for_prefill(cls, - final_state: torch.Tensor, - raw_states: list[torch.Tensor], - prefill_metadata: ConceptPrefillMetadata): - """Build decoder concept inputs from compact prefill predictor - states.""" - return cls( - final_state=final_state, - raw_states=raw_states, - prefill_metadata=prefill_metadata, - ) - - @classmethod - def for_decode(cls, final_state: torch.Tensor, decode_states: torch.Tensor): - """Build decoder concept inputs from already gathered decode states.""" - return cls( - final_state=final_state, - decode_states=decode_states, - ) - - @dataclass class _EncoderOutput: """Encoder result plus its reusable layer-major SelfDD history buffer.""" @@ -198,22 +167,15 @@ def forward(self, else: hidden_states = inputs_embeds - if concept_metadata.is_decoding: - return self._forward_decode( - hidden_states, - position_ids, - concept_metadata, - concept_caches, - ) - - hidden_states, prefill_position_ids = self._normalize_prefill_inputs( + hidden_states, position_ids = self._normalize_forward_inputs( hidden_states, position_ids, - attn_metadata, + concept_metadata, ) - return self._forward_prefill_packed( + self._validate_concept_caches(concept_metadata, concept_caches) + return self._forward_token_stream( hidden_states, - prefill_position_ids, + position_ids, concept_metadata, concept_caches, ) @@ -250,27 +212,15 @@ def _build_concept_caches(self, named_state_caches=named_state_caches, ) - def _decode_chunk_state_update(self, - current_source_states: torch.Tensor, - concept_metadata: ConceptMetadata, - concept_caches: ConceptCaches) -> ConceptChunkStateUpdateResult: - """Update decode chunk-source state and return fixed-shape rows. - - CUDA uses the Triton writer. CPU uses the reference writer for tests. The returned rows deliberately avoid - dynamic concept-row compaction, matching the CUDA graph route in the design doc. - """ - chunk_source_state = concept_caches.chunk_source_state - concept_input_states, update_mask = self.concept_ops.decode_chunk_state_update( - chunk_source_state, - current_source_states, - concept_metadata.state_ids, - concept_metadata.position_ids, - concept_metadata.chunk_size, - concept_metadata.merge_method, - ) - return ConceptChunkStateUpdateResult( - concept_input_states=concept_input_states, - concept_update_mask=update_mask, + @staticmethod + def _build_runtime_caches(concept_caches: ConceptCaches) -> ConceptRuntimeCaches: + """Expose only backend-owned cache views to ConceptLM runtime ops.""" + return ConceptRuntimeCaches( + chunk_source_state=concept_caches.chunk_source_state, + last_state=concept_caches.last_state, + last_raw_states=concept_caches.last_raw_states, + last_final_state=concept_caches.last_final_state, + concept_past_key_values=concept_caches.concept_past_key_values, ) def _route_gate(self, layer_idx: int) -> torch.Tensor: @@ -278,6 +228,23 @@ def _route_gate(self, layer_idx: int) -> torch.Tensor: concept_route_scale]``.""" return self.final_read_concept_gate_logits[int(layer_idx)].float().softmax(dim=-1) + def _normalize_forward_inputs(self, + hidden_states: torch.Tensor, + position_ids: torch.Tensor, + concept_metadata: ConceptMetadata): + """Normalize token-stream inputs to flat ``[tokens_or_batch, hidden]``. + + Prefill and decode use different engine layouts, but the model body below consumes one flat token stream for + both phases. + """ + if concept_metadata.is_decoding: + return self._normalize_decode_inputs(hidden_states, position_ids) + return self._normalize_prefill_inputs( + hidden_states, + position_ids, + concept_metadata.attn_metadata, + ) + def _normalize_prefill_inputs(self, hidden_states: torch.Tensor, position_ids: torch.Tensor, @@ -327,151 +294,44 @@ def _normalize_decode_inputs(self, hidden_states: torch.Tensor, position_ids: to raise ValueError(f'Expected {batch_size} decode position ids, got {position_ids.numel()}.') return hidden_states[0].contiguous(), position_ids - def _select_decode_last_state_rows(self, concept_caches: ConceptCaches, - decode_metadata: ConceptDecodeMetadata) -> tuple[torch.Tensor, torch.Tensor]: - """Gather latest concept state rows through backend-owned layout.""" - return self.concept_ops.select_decode_last_state_rows( - concept_caches.last_state, - concept_caches.last_final_state, - concept_caches.last_raw_states, - decode_metadata, - ) - - def _select_decode_decoder_concepts(self, - concept_caches: ConceptCaches, - decode_metadata: ConceptDecodeMetadata, - previous_final_state: torch.Tensor | None, - previous_raw_states: torch.Tensor | None) -> _DecoderConceptInput: - """Select the final/raw concept state visible to this decode token.""" - if bool(getattr(self.config, 'concept_shift_feature', True)): - final_state, raw_states = self._select_decode_last_state_rows(concept_caches, decode_metadata) - else: - final_state = previous_final_state - raw_states = previous_raw_states - - concept_read_mask = self.concept_ops.decode_concept_read_mask(decode_metadata) - final_state = torch.where(concept_read_mask.view(-1, 1), final_state, torch.zeros_like(final_state)) - raw_states = torch.where(concept_read_mask.view(-1, 1, 1), raw_states, torch.zeros_like(raw_states)) - return _DecoderConceptInput.for_decode(final_state, raw_states) + @staticmethod + def _validate_concept_caches(concept_metadata: ConceptMetadata, concept_caches: ConceptCaches) -> None: + """Validate cache streams needed by the shared forward body.""" + phase = 'decode' if concept_metadata.is_decoding else 'prefill' + if (concept_caches.encoder_past_key_values is None or concept_caches.concept_past_key_values is None + or concept_caches.decoder_past_key_values is None): + raise RuntimeError(f'ConceptLM {phase} requires encoder, concept, and decoder KV caches.') + if not concept_metadata.is_decoding: + return + if concept_caches.chunk_source_state is None: + raise RuntimeError('ConceptLM decode requires chunk source state cache.') + if concept_caches.last_final_state is None or concept_caches.last_raw_states is None: + raise RuntimeError('ConceptLM decode requires cached last concept states.') @staticmethod - def _build_decode_chunk_source_states(encoder_output: _EncoderOutput) -> torch.Tensor: - """Return current per-row states accumulated until the next concept - boundary as a view over encoder history. + def _build_chunk_source_states(encoder_output: _EncoderOutput) -> torch.Tensor: + """Return source states consumed by concept chunk preparation. Row 0 is the final encoder hidden, used as concept-predictor input when - a boundary is reached. Remaining rows mirror the prefill - ``encoder_raw_states[:-1]`` route sources. + a chunk concept is produced. Remaining rows mirror + ``encoder_raw_states[:-1]`` route sources. The returned layout is shared + by prefill and decode: ``[token_or_batch, num_sources, hidden]``. """ num_sources = max(len(encoder_output.raw_states), 1) return encoder_output.history_buffer[:num_sources].movedim(0, -2) - @staticmethod - def _decode_concept_position_ids(position_ids: torch.Tensor, chunk_size: int) -> torch.Tensor: - """Return reference RoPE positions for concept rows emitted at decode - boundaries.""" - return (position_ids - int(chunk_size) + 1).clamp(min=0) - - def _snapshot_decode_concept_kv(self, - concept_caches: ConceptCaches, - concept_attn_metadata: Any): - """Snapshot concept KV slots that dummy non-boundary rows may - overwrite.""" - return self.concept_ops.snapshot_decode_concept_kv( - concept_caches.concept_past_key_values, - concept_attn_metadata, - ) - - def _restore_decode_concept_kv_( - self, - concept_caches: ConceptCaches, - concept_attn_metadata: Any, - saved_kv, - restore_mask: torch.Tensor, - ): - """Restore concept KV slots for non-boundary and padded rows.""" - self.concept_ops.restore_decode_concept_kv( - concept_caches.concept_past_key_values, - concept_attn_metadata, - saved_kv, - restore_mask, - ) - - def _write_decode_concept_states_static_(self, - concept_caches: ConceptCaches, - decode_metadata: ConceptDecodeMetadata, - update_mask: torch.Tensor, - predicted_vectors: torch.Tensor, - concept_raw_states: list[torch.Tensor]): - """Write newly emitted concept states to persistent decode caches.""" - self.concept_ops.write_decode_concept_states( - concept_caches.last_raw_states, - concept_caches.last_final_state, - predicted_vectors, - concept_raw_states, - decode_metadata.state_ids, - update_mask, - ) - - def _build_decode_concept_request(self, - chunk_update: ConceptChunkStateUpdateResult, - decode_metadata: ConceptDecodeMetadata, - concept_metadata: ConceptMetadata) -> _ConceptPredictorRequest: - """Build fixed-shape concept predictor inputs for decode.""" - concept_hidden = self.concept_vq_input_norm(chunk_update.concept_input_states[:, 0]) - encoder_concept_states = self.concept_predictor.normalize_encoder_concept_states( - chunk_update.concept_input_states[:, 1:]) - concept_position_ids = self._decode_concept_position_ids( - decode_metadata.position_ids, - concept_metadata.chunk_size, - ) - concept_attn_metadata = self.concept_ops.build_concept_decode_metadata_static( - concept_metadata.attn_metadata, - decode_metadata, - ) + def _build_concept_request(self, chunk_input: ConceptChunkInput) -> _ConceptPredictorRequest: + """Build concept predictor inputs from backend-prepared chunk rows.""" + concept_hidden = self.concept_vq_input_norm(chunk_input.source_states[:, 0]) + encoder_sources = chunk_input.source_states[:, 1:] + encoder_concept_states = self.concept_predictor.normalize_encoder_concept_states(encoder_sources) return _ConceptPredictorRequest( hidden_states=concept_hidden, encoder_states=encoder_concept_states, - position_ids=concept_position_ids, - attn_metadata=concept_attn_metadata, - ) - - def _update_decode_concept_states_static_(self, - chunk_update: ConceptChunkStateUpdateResult, - decode_metadata: ConceptDecodeMetadata, - concept_metadata: ConceptMetadata, - concept_caches: ConceptCaches): - """Emit/cache concept states with fixed batch shape. - - The predictor runs for every decode row so CUDA graph capture sees a stable launch sequence. Non-boundary rows - are dummy work: their concept KV writes are restored and their final/raw state writes are masked. - """ - request = self._build_decode_concept_request(chunk_update, decode_metadata, concept_metadata) - saved_kv = self._snapshot_decode_concept_kv(concept_caches, request.attn_metadata) - concept_output = self._run_concept_predictor(request, concept_caches) - self._restore_decode_concept_kv_( - concept_caches, - request.attn_metadata, - saved_kv, - ~chunk_update.concept_update_mask, - ) - self._write_decode_concept_states_static_( - concept_caches, - decode_metadata, - chunk_update.concept_update_mask, - concept_output.predicted_vectors, - concept_output.raw_states, + position_ids=chunk_input.position_ids, + attn_metadata=chunk_input.attn_metadata, ) - def _build_encoder_concept_states_packed(self, - encoder_raw_states: list[torch.Tensor], - prefill_metadata: ConceptPrefillMetadata) -> torch.Tensor: - """Build packed chunk-level encoder states used by the concept - predictor.""" - chunks = [self.concept_ops.merge_chunks_packed(state, prefill_metadata) for state in encoder_raw_states[:-1]] - states = torch.stack(chunks, dim=-2) - return self.concept_predictor.normalize_encoder_concept_states(states) - def _run_concept_predictor(self, request: _ConceptPredictorRequest, concept_caches: ConceptCaches) -> _ConceptPredictorOutput: @@ -489,26 +349,6 @@ def _run_concept_predictor(self, raw_states=concept_raw_states, ) - def _build_prefill_concept_request(self, - hidden_states: torch.Tensor, - encoder_raw_states: list[torch.Tensor], - prefill_metadata: ConceptPrefillMetadata, - concept_metadata: ConceptMetadata) -> _ConceptPredictorRequest: - """Build concept predictor inputs for compact packed prefill.""" - concept_hidden = self.concept_ops.merge_chunks_packed(hidden_states, prefill_metadata) - concept_hidden = self.concept_vq_input_norm(concept_hidden) - encoder_concept_states = self._build_encoder_concept_states_packed(encoder_raw_states, prefill_metadata) - concept_attn_metadata = self.concept_ops.build_concept_prefill_metadata( - concept_metadata.attn_metadata, - prefill_metadata, - ) - return _ConceptPredictorRequest( - hidden_states=concept_hidden, - encoder_states=encoder_concept_states, - position_ids=prefill_metadata.concept_position_ids, - attn_metadata=concept_attn_metadata, - ) - def _fuse_token_concept_states(self, hidden_states: torch.Tensor, final_concept_state: torch.Tensor) -> torch.Tensor: @@ -519,7 +359,7 @@ def _fuse_token_concept_states(self, def _run_decoder_from_concepts(self, hidden_states: torch.Tensor, encoder_raw_states: list[torch.Tensor], - decoder_concepts: _DecoderConceptInput, + decoder_concepts: ConceptDecoderInput, position_ids: torch.Tensor, concept_caches: ConceptCaches, attn_metadata: Any) -> torch.Tensor: @@ -536,34 +376,21 @@ def _run_decoder_from_concepts(self, def _prepare_decoder_route_sources(self, encoder_raw_states: list[torch.Tensor], - decoder_concepts: _DecoderConceptInput) -> tuple[torch.Tensor, torch.Tensor]: + decoder_concepts: ConceptDecoderInput) -> tuple[torch.Tensor, torch.Tensor]: """Prepare encoder/concept route sources in packed ``[..., L, H]`` layout.""" decoder_encoder_states = torch.stack(tuple(encoder_raw_states), dim=-2) decoder_encoder_states = self.decoder_read_encoder_shared_source_norm(decoder_encoder_states) - if decoder_concepts.decode_states is not None: - concept_states = self.decoder_read_concept_shared_source_norm(decoder_concepts.decode_states) - return decoder_encoder_states, concept_states - - concept_states = torch.stack(tuple(decoder_concepts.raw_states), dim=-2) - zero_chunk = torch.zeros_like(concept_states[:1]) - concept_states = torch.cat((zero_chunk, concept_states), dim=0) - concept_states = self.decoder_read_concept_shared_source_norm(concept_states) - concept_states = self.concept_ops.repeat_shift_source_states_packed(concept_states, - decoder_concepts.prefill_metadata) + concept_states = self.decoder_read_concept_shared_source_norm(decoder_concepts.route_states) return decoder_encoder_states, concept_states - def _forward_prefill_packed(self, - hidden_states: torch.Tensor, - position_ids: torch.Tensor, - concept_metadata: ConceptMetadata, - concept_caches: ConceptCaches): - """Packed non-decode ConceptLM forward.""" - if (concept_caches.encoder_past_key_values is None or concept_caches.concept_past_key_values is None - or concept_caches.decoder_past_key_values is None): - raise RuntimeError('ConceptLM prefill requires encoder, concept, and decoder KV caches.') - prefill_metadata = self.concept_ops.build_prefill_metadata(concept_metadata.attn_metadata, position_ids) + def _forward_token_stream(self, + hidden_states: torch.Tensor, + position_ids: torch.Tensor, + concept_metadata: ConceptMetadata, + concept_caches: ConceptCaches): + """Shared ConceptLM forward for normalized prefill and decode rows.""" encoder_output = self._encode( hidden_states, position_ids, @@ -572,28 +399,32 @@ def _forward_prefill_packed(self, ) hidden_states = encoder_output.hidden_states encoder_raw_states = encoder_output.raw_states - concept_request = self._build_prefill_concept_request( - hidden_states, - encoder_raw_states, - prefill_metadata, - concept_metadata, + source_states = self._build_chunk_source_states(encoder_output) + runtime_caches = self._build_runtime_caches(concept_caches) + chunk_input = self.concept_ops.build_concept_chunk_input( + source_states, + concept_metadata.attn_metadata, + position_ids, + state_ids=concept_metadata.state_ids, + chunk_source_state_cache=runtime_caches.chunk_source_state, ) + forward_context = self.concept_ops.begin_concept_forward(chunk_input, runtime_caches) + concept_request = self._build_concept_request(chunk_input) concept_output = self._run_concept_predictor(concept_request, concept_caches) - self.concept_ops.write_prefill_state_caches_eager( - concept_caches.chunk_source_state, - concept_caches.last_raw_states, - concept_caches.last_final_state, - concept_metadata.state_ids, - prefill_metadata, - self._build_decode_chunk_source_states(encoder_output), + self.concept_ops.end_concept_forward( + chunk_input, + runtime_caches, + forward_context, + source_states, concept_output.predicted_vectors, concept_output.raw_states, ) - repeated_concepts = self.concept_ops.repeat_shift_packed(concept_output.predicted_vectors, prefill_metadata) - decoder_concepts = _DecoderConceptInput.for_prefill( - repeated_concepts, + decoder_concepts = self.concept_ops.build_decoder_concept_input( + chunk_input, + runtime_caches, + forward_context, + concept_output.predicted_vectors, concept_output.raw_states, - prefill_metadata, ) final_hidden = self._run_decoder_from_concepts( hidden_states, @@ -605,80 +436,6 @@ def _forward_prefill_packed(self, ) return final_hidden.unsqueeze(0).contiguous() - def _forward_decode(self, - hidden_states: torch.Tensor, - position_ids: torch.Tensor, - concept_metadata: ConceptMetadata, - concept_caches: ConceptCaches): - """ConceptLM decode path. - - This path is semantically structured for serving. Boundary concept - updates run with fixed batch shape so it is eligible for CUDA graph - replay through the base ``CudaGraphMixin`` decode-only policy. - """ - if (concept_caches.encoder_past_key_values is None or concept_caches.concept_past_key_values is None - or concept_caches.decoder_past_key_values is None): - raise RuntimeError('ConceptLM decode requires encoder, concept, and decoder KV caches.') - if concept_caches.chunk_source_state is None: - raise RuntimeError('ConceptLM decode requires chunk source state cache.') - if concept_caches.last_final_state is None or concept_caches.last_raw_states is None: - raise RuntimeError('ConceptLM decode requires cached last concept states.') - - hidden_states, decode_position_ids = self._normalize_decode_inputs(hidden_states, position_ids) - decode_metadata = self.concept_ops.build_decode_metadata( - decode_position_ids, - concept_metadata.state_ids, - hidden_states.size(0), - hidden_states.device, - ) - decode_concept_metadata = replace( - concept_metadata, - position_ids=decode_position_ids, - state_ids=decode_metadata.state_ids, - ) - - encoder_output = self._encode( - hidden_states, - decode_position_ids, - past_key_values=concept_caches.encoder_past_key_values, - attn_metadata=concept_metadata.attn_metadata, - ) - hidden_states = encoder_output.hidden_states - encoder_raw_states = encoder_output.raw_states - previous_final_concept_state = None - previous_concept_raw_state_rows = None - if not bool(getattr(self.config, 'concept_shift_feature', True)): - previous_final_concept_state, previous_concept_raw_state_rows = self._select_decode_last_state_rows( - concept_caches, decode_metadata) - current_source_states = self._build_decode_chunk_source_states(encoder_output) - chunk_update = self._decode_chunk_state_update( - current_source_states, - decode_concept_metadata, - concept_caches, - ) - self._update_decode_concept_states_static_( - chunk_update, - decode_metadata, - decode_concept_metadata, - concept_caches, - ) - - decoder_concepts = self._select_decode_decoder_concepts( - concept_caches, - decode_metadata, - previous_final_concept_state, - previous_concept_raw_state_rows, - ) - final_hidden = self._run_decoder_from_concepts( - hidden_states, - encoder_raw_states, - decoder_concepts, - decode_position_ids, - concept_caches, - concept_metadata.attn_metadata, - ) - return final_hidden.unsqueeze(0).contiguous() - def _encode(self, hidden_states: torch.Tensor, position_ids: torch.Tensor, @@ -721,7 +478,7 @@ def _encode(self, def _decode(self, decoder_input: torch.Tensor, encoder_raw_states: list[torch.Tensor], - decoder_concepts: _DecoderConceptInput, + decoder_concepts: ConceptDecoderInput, position_ids: torch.Tensor, concept_caches: ConceptCaches, attn_metadata: Any = None): diff --git a/lmdeploy/pytorch/nn/conceptlm.py b/lmdeploy/pytorch/nn/conceptlm.py index 9ce2c60cff..c4098ffa8f 100644 --- a/lmdeploy/pytorch/nn/conceptlm.py +++ b/lmdeploy/pytorch/nn/conceptlm.py @@ -5,7 +5,12 @@ from torch import Tensor, nn from lmdeploy.pytorch.backends import OpType, get_backend -from lmdeploy.pytorch.backends.conceptlm import ConceptDecodeMetadata, ConceptPrefillMetadata +from lmdeploy.pytorch.backends.conceptlm import ( + ConceptChunkInput, + ConceptDecoderInput, + ConceptForwardContext, + ConceptRuntimeCaches, +) class ConceptLMRuntimeOps(nn.Module): @@ -26,189 +31,60 @@ def flatten_decode_position_ids(self, position_ids: Tensor, batch_size: int, dev row.""" return self.impl.flatten_decode_position_ids(position_ids, batch_size, device) - def build_decode_metadata(self, - position_ids: Tensor, - state_ids: Tensor | None, - batch_size: int, - device: torch.device) -> ConceptDecodeMetadata: - """Build fixed-shape decode metadata from engine state ids.""" - return self.impl.build_decode_metadata(position_ids, state_ids, batch_size, device) - - def select_decode_last_state_rows(self, - last_state: Tensor | None, - last_final_state: Tensor, - last_raw_states: Tensor, - decode_metadata: ConceptDecodeMetadata) -> tuple[Tensor, Tensor]: - """Gather latest concept state rows for decode.""" - return self.impl.select_decode_last_state_rows( - last_state, - last_final_state, - last_raw_states, - decode_metadata, + def build_concept_chunk_input( + self, + source_states: Tensor, + token_attn_metadata: Any, + position_ids: Tensor, + state_ids: Tensor | None = None, + chunk_source_state_cache: Tensor | None = None, + ) -> ConceptChunkInput: + """Build concept-predictor source rows for prefill or decode.""" + return self.impl.build_concept_chunk_input( + source_states, + token_attn_metadata, + position_ids, + state_ids=state_ids, + chunk_source_state_cache=chunk_source_state_cache, ) - def decode_concept_read_mask(self, decode_metadata: ConceptDecodeMetadata) -> Tensor: - """Return rows whose current decode token should read a cached - concept.""" - return self.impl.decode_concept_read_mask(decode_metadata) - - def build_concept_decode_metadata_static(self, token_attn_metadata: Any, - decode_metadata: ConceptDecodeMetadata): - """Build fixed-shape concept-stream decode metadata.""" - return self.impl.build_concept_decode_metadata_static(token_attn_metadata, decode_metadata) - - def build_prefill_metadata(self, token_attn_metadata: Any, position_ids: Tensor) -> ConceptPrefillMetadata: - """Build packed token-to-concept metadata for batched prefill.""" - return self.impl.build_prefill_metadata(token_attn_metadata, position_ids) - - def merge_chunks_packed(self, hidden_states: Tensor, prefill_metadata: ConceptPrefillMetadata) -> Tensor: - """Merge packed token states into compact concept rows.""" - return self.impl.merge_chunks_packed(hidden_states, prefill_metadata) - - def repeat_shift_packed(self, concept_states: Tensor, prefill_metadata: ConceptPrefillMetadata) -> Tensor: - """Gather compact concept states back to packed token rows.""" - return self.impl.repeat_shift_packed(concept_states, prefill_metadata) + def begin_concept_forward(self, chunk_input: ConceptChunkInput, + runtime_caches: ConceptRuntimeCaches) -> ConceptForwardContext: + """Prepare transient state before the concept predictor forward.""" + return self.impl.begin_concept_forward(chunk_input, runtime_caches) - def repeat_shift_source_states_packed(self, concept_states_with_zero: Tensor, - prefill_metadata: ConceptPrefillMetadata) -> Tensor: - """Gather zero-prefixed compact concept route states to token rows.""" - return self.impl.repeat_shift_source_states_packed(concept_states_with_zero, prefill_metadata) - - def build_concept_prefill_metadata(self, token_attn_metadata: Any, prefill_metadata: ConceptPrefillMetadata): - """Build chunk-stream attention metadata for packed prefill.""" - return self.impl.build_concept_prefill_metadata(token_attn_metadata, prefill_metadata) - - def write_prefill_state_caches_eager( + def end_concept_forward( self, - chunk_source_state: Tensor | None, - last_raw_states: Tensor | None, - last_final_state: Tensor | None, - state_ids: Tensor | None, - prefill_metadata: ConceptPrefillMetadata, + chunk_input: ConceptChunkInput, + runtime_caches: ConceptRuntimeCaches, + forward_context: ConceptForwardContext, source_states: Tensor, predicted_vectors: Tensor, concept_raw_states: list[Tensor], ) -> None: - """Seed decode state caches from a completed prefill forward.""" - return self.impl.write_prefill_state_caches_eager( - chunk_source_state, - last_raw_states, - last_final_state, - state_ids, - prefill_metadata, + """Commit concept-predictor side effects for prefill or decode.""" + return self.impl.end_concept_forward( + chunk_input, + runtime_caches, + forward_context, source_states, predicted_vectors, concept_raw_states, ) - def stack_concept_raw_states(self, concept_raw_states: list[Tensor]) -> Tensor: - """Stack raw concept-layer states.""" - return self.impl.stack_concept_raw_states(concept_raw_states) - - def snapshot_decode_concept_kv(self, concept_past_key_values: list[list[Tensor]], - concept_attn_metadata: Any) -> list[tuple[Tensor, Tensor]]: - """Snapshot concept KV slots that dummy rows may overwrite.""" - return self.impl.snapshot_decode_concept_kv(concept_past_key_values, concept_attn_metadata) - - def restore_decode_concept_kv(self, concept_past_key_values: list[list[Tensor]], concept_attn_metadata: Any, - saved_kv: list[tuple[Tensor, Tensor]], restore_mask: Tensor) -> None: - """Restore concept KV slots for dummy rows.""" - return self.impl.restore_decode_concept_kv( - concept_past_key_values, - concept_attn_metadata, - saved_kv, - restore_mask, - ) - - def write_decode_concept_states( + def build_decoder_concept_input( self, - last_raw_state_cache: Tensor, - last_final_state_cache: Tensor, + chunk_input: ConceptChunkInput, + runtime_caches: ConceptRuntimeCaches, + forward_context: ConceptForwardContext, predicted_vectors: Tensor, - raw_states: list[Tensor], - state_ids: Tensor, - update_mask: Tensor, - ) -> None: - """Write newly emitted decode concept states.""" - return self.impl.write_decode_concept_states( - last_raw_state_cache, - last_final_state_cache, - predicted_vectors, - raw_states, - state_ids, - update_mask, - ) - - def decode_chunk_state_update( - self, - chunk_source_state_cache: Tensor, - current_source_states: Tensor, - state_ids: Tensor, - position_ids: Tensor, - chunk_size: int, - merge_method: str, - ) -> tuple[Tensor, Tensor]: - """Update state cache and return concept inputs plus update mask.""" - return self.impl.decode_chunk_state_update( - chunk_source_state_cache, - current_source_states, - state_ids, - position_ids, - chunk_size, - merge_method, - ) - - def decode_kv_cache_snapshot( - self, - k_cache: Tensor, - v_cache: Tensor, - block_offsets: Tensor, - kv_seqlens: Tensor, - ) -> tuple[Tensor, Tensor]: - """Snapshot one decode KV slot per batch row.""" - return self.impl.decode_kv_cache_snapshot( - k_cache, - v_cache, - block_offsets, - kv_seqlens, - ) - - def decode_kv_cache_restore( - self, - k_cache: Tensor, - v_cache: Tensor, - saved_k: Tensor, - saved_v: Tensor, - block_offsets: Tensor, - kv_seqlens: Tensor, - restore_mask: Tensor, - ) -> None: - """Restore one decode KV slot for masked batch rows.""" - return self.impl.decode_kv_cache_restore( - k_cache, - v_cache, - saved_k, - saved_v, - block_offsets, - kv_seqlens, - restore_mask, - ) - - def decode_concept_state_update( - self, - last_raw_state_cache: Tensor, - last_final_state_cache: Tensor, - predicted_vectors: Tensor, - raw_states: Tensor, - state_ids: Tensor, - update_mask: Tensor, - ) -> None: - """Write final/raw concept states for masked decode rows.""" - return self.impl.decode_concept_state_update( - last_raw_state_cache, - last_final_state_cache, + concept_raw_states: list[Tensor], + ) -> ConceptDecoderInput: + """Build token-decoder concept inputs for prefill or decode.""" + return self.impl.build_decoder_concept_input( + chunk_input, + runtime_caches, + forward_context, predicted_vectors, - raw_states, - state_ids, - update_mask, + concept_raw_states, ) From 7e04f921f4998c96cd6530959e3f1b9ff0f9671a Mon Sep 17 00:00:00 2001 From: grimoire Date: Sun, 26 Jul 2026 18:09:33 +0800 Subject: [PATCH 10/16] Rename ConceptLM prefill cache writer hook --- lmdeploy/pytorch/backends/cuda/conceptlm.py | 4 ++-- lmdeploy/pytorch/backends/default/conceptlm.py | 4 ++-- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/lmdeploy/pytorch/backends/cuda/conceptlm.py b/lmdeploy/pytorch/backends/cuda/conceptlm.py index c0749548ee..64826867e0 100644 --- a/lmdeploy/pytorch/backends/cuda/conceptlm.py +++ b/lmdeploy/pytorch/backends/cuda/conceptlm.py @@ -47,7 +47,7 @@ def prefill_chunk_state_update(self, source_states: Tensor, prefill_metadata) -> self.merge_method, ) - def write_prefill_state_caches_eager( + def _write_prefill_state_caches_impl( self, chunk_source_state: Tensor | None, last_raw_states: Tensor | None, @@ -65,7 +65,7 @@ def write_prefill_state_caches_eager( return if not (source_states.is_cuda and chunk_source_state.is_cuda and last_raw_states.is_cuda and last_final_state.is_cuda): - return super().write_prefill_state_caches_eager( + return super()._write_prefill_state_caches_impl( chunk_source_state, last_raw_states, last_final_state, diff --git a/lmdeploy/pytorch/backends/default/conceptlm.py b/lmdeploy/pytorch/backends/default/conceptlm.py index d307329658..134fa64e8f 100644 --- a/lmdeploy/pytorch/backends/default/conceptlm.py +++ b/lmdeploy/pytorch/backends/default/conceptlm.py @@ -673,7 +673,7 @@ def write_prefill_state_caches( concept_raw_states: list[Tensor], ) -> None: """Seed decode state caches from a completed prefill forward.""" - return self.write_prefill_state_caches_eager( + return self._write_prefill_state_caches_impl( chunk_source_state, last_raw_states, last_final_state, @@ -684,7 +684,7 @@ def write_prefill_state_caches( concept_raw_states, ) - def write_prefill_state_caches_eager( + def _write_prefill_state_caches_impl( self, chunk_source_state: Tensor | None, last_raw_states: Tensor | None, From b5497e381e02e2e7cb99ca2c33007b088cd8dbd4 Mon Sep 17 00:00:00 2001 From: grimoire Date: Tue, 28 Jul 2026 20:01:36 +0800 Subject: [PATCH 11/16] Fix ConceptLM decode state alignment --- .../pytorch/backends/default/conceptlm.py | 22 +++++++++++++------ lmdeploy/pytorch/configurations/conceptlm.py | 18 ++++++++++----- lmdeploy/pytorch/models/intern_ncp/modules.py | 4 ++-- 3 files changed, 29 insertions(+), 15 deletions(-) diff --git a/lmdeploy/pytorch/backends/default/conceptlm.py b/lmdeploy/pytorch/backends/default/conceptlm.py index 134fa64e8f..2cc2744353 100644 --- a/lmdeploy/pytorch/backends/default/conceptlm.py +++ b/lmdeploy/pytorch/backends/default/conceptlm.py @@ -631,7 +631,8 @@ def merge_prefill_tail_chunk_states(self, source_states: Tensor, tail_lens = torch.where(q_seqlens > 0, tail_lens, torch.zeros_like(tail_lens)) has_tail = tail_lens > 0 - tail_rows = source_states.new_zeros((batch_size, source_states.size(1), source_states.size(2))) + tail_rows = source_states.new_zeros((batch_size, source_states.size(1), source_states.size(2)), + dtype=torch.float32) if source_states.size(0) == 0: return tail_rows if self.merge_method == 'first': @@ -651,7 +652,7 @@ def merge_prefill_tail_chunk_states(self, source_states: Tensor, token_pos = token_ids - cu_q_seqlens[token_seq] token_tail_start = q_seqlens[token_seq] - tail_lens[token_seq] valid_tail = (tail_lens[token_seq] > 0) & (token_pos >= token_tail_start) - weighted_source = source_states * valid_tail.to(dtype=source_states.dtype).view(-1, 1, 1) + weighted_source = source_states.float() * valid_tail.to(dtype=torch.float32).view(-1, 1, 1) tail_rows.index_add_(0, token_seq, weighted_source) return tail_rows @@ -814,7 +815,9 @@ def decode_chunk_state_update( valid_state_mask = state_ids >= 0 safe_state_ids = state_ids.clamp(min=0) + accumulator_dtype = chunk_source_state_cache.dtype previous_rows = chunk_source_state_cache.index_select(0, safe_state_ids) + current_rows = current_source_states.to(dtype=accumulator_dtype) chunk_size = int(chunk_size) chunk_pos = torch.remainder(position_ids, chunk_size) @@ -823,19 +826,24 @@ def decode_chunk_state_update( merge_method = str(merge_method) if merge_method == 'first': - update_rows = torch.where(first_token_mask.view(batch_size, 1, 1), current_source_states, previous_rows) + update_rows = torch.where(first_token_mask.view(batch_size, 1, 1), current_rows, previous_rows) concept_input_states = update_rows elif merge_method == 'last': - update_rows = current_source_states - concept_input_states = current_source_states + update_rows = current_rows + concept_input_states = current_rows else: - update_rows = previous_rows + current_source_states + update_rows = previous_rows + current_rows concept_input_states = update_rows / chunk_size zero_rows = torch.zeros_like(update_rows) next_rows = torch.where(update_mask.view(batch_size, 1, 1), zero_rows, update_rows) next_rows = torch.where(valid_state_mask.view(batch_size, 1, 1), next_rows, previous_rows) - concept_input_states = torch.where(update_mask.view(batch_size, 1, 1), concept_input_states, zero_rows) + concept_zero_rows = torch.zeros_like(current_source_states) + concept_input_states = torch.where( + update_mask.view(batch_size, 1, 1), + concept_input_states.to(dtype=current_source_states.dtype), + concept_zero_rows, + ) for batch_idx in range(batch_size): state_id = int(state_ids[batch_idx]) diff --git a/lmdeploy/pytorch/configurations/conceptlm.py b/lmdeploy/pytorch/configurations/conceptlm.py index 5ce5b75ee3..4ba7d64e5b 100644 --- a/lmdeploy/pytorch/configurations/conceptlm.py +++ b/lmdeploy/pytorch/configurations/conceptlm.py @@ -68,12 +68,18 @@ def build(cls, hf_config, model_path: str = None, **kwargs): model_config.llm_config.concept_kv_total_layers = model_config.num_layers hidden_size = int(hf_config.hidden_size) - state_dtype = _get_concept_state_dtype(hf_config) + last_state_dtype = _get_concept_state_dtype(hf_config) concept_encoder_read_sources = max(enc_layers - 1, 0) # Decode accumulates the current chunk for every state needed when a - # chunk boundary emits one concept. Row 0 is the final encoder hidden - # used as concept-predictor input; following rows are encoder raw states - # consumed by concept-read-encoder residual routes. + # chunk boundary emits one concept. Keep this accumulator in fp32 to + # match the reference chunk merge: reduce the whole chunk, then cast the + # emitted concept input back to model dtype. Using bf16 here would round + # the partial sum after every decode token and drift from full-forward + # semantics. + # + # Row 0 is the final encoder hidden used as concept-predictor input; + # following rows are encoder raw states consumed by concept-read-encoder + # residual routes. concept_chunk_state_sources = 1 + concept_encoder_read_sources # Last emitted concept snapshot. Row 0 is the final concept vector, # rows 1: are raw concept-layer states. Keep this packed so decode can @@ -81,9 +87,9 @@ def build(cls, hf_config, model_path: str = None, **kwargs): concept_last_state_sources = 1 + concept_layers state_specs = [ StateCacheSpec(CONCEPT_STATE_NAMES[CONCEPT_STATE_CHUNK_SOURCE], - (concept_chunk_state_sources, hidden_size), state_dtype), + (concept_chunk_state_sources, hidden_size), torch.float32), StateCacheSpec(CONCEPT_STATE_NAMES[CONCEPT_STATE_LAST], (concept_last_state_sources, hidden_size), - state_dtype), + last_state_dtype), ] model_config.state_cache_specs = state_specs # Backward-compat bridge used by scheduler/state-cache sizing. The diff --git a/lmdeploy/pytorch/models/intern_ncp/modules.py b/lmdeploy/pytorch/models/intern_ncp/modules.py index 7e77eae76f..da16de9d31 100644 --- a/lmdeploy/pytorch/models/intern_ncp/modules.py +++ b/lmdeploy/pytorch/models/intern_ncp/modules.py @@ -877,9 +877,9 @@ def __init__(self, @staticmethod def _layer_sliding_window(layer_number: int, window_size: int | None, skip_frequency: int | None): """Match reference OLMo window/full attention alternation.""" - if window_size is None: + if window_size is None or skip_frequency is None: return None - if skip_frequency is not None and layer_number % skip_frequency == 0: + if layer_number % skip_frequency == 0: return None return window_size From 1a305a4f327c7cd92c31b80ece813d589286c951 Mon Sep 17 00:00:00 2001 From: grimoire Date: Wed, 29 Jul 2026 11:14:57 +0800 Subject: [PATCH 12/16] Fix ConceptLM compressed RoPE positions --- .../pytorch/backends/default/conceptlm.py | 20 ++++++- .../backends/test_conceptlm_runtime_ops.py | 60 +++++++++++++++++++ 2 files changed, 77 insertions(+), 3 deletions(-) create mode 100644 tests/pytorch/backends/test_conceptlm_runtime_ops.py diff --git a/lmdeploy/pytorch/backends/default/conceptlm.py b/lmdeploy/pytorch/backends/default/conceptlm.py index 2cc2744353..2253c2d8c8 100644 --- a/lmdeploy/pytorch/backends/default/conceptlm.py +++ b/lmdeploy/pytorch/backends/default/conceptlm.py @@ -263,8 +263,13 @@ def build_concept_chunk_input( ) def decode_concept_position_ids(self, position_ids: Tensor) -> Tensor: - """Return reference RoPE positions for emitted decode concept rows.""" - return (position_ids - self.chunk_size + 1).clamp(min=0) + """Return compressed-timeline RoPE positions for decode concept rows.""" + concept_index = torch.div( + position_ids + 1, + self.chunk_size, + rounding_mode='floor', + ) - 1 + return concept_index.clamp(min=0) def _get_max_concepts_per_request(self, token_attn_metadata: Any, concept_q_seqlens: Tensor) -> int: """Return per-request concept attention bound without hidden context @@ -327,7 +332,16 @@ def _build_prefill_concept_layout(self, token_attn_metadata: Any, position_ids: concept_seq = torch.searchsorted(concept_cu_seqlens_long[1:], concept_ids, right=True) local_concept_ids = concept_ids - concept_cu_seqlens_long[concept_seq] concept_token_start = token_layout.q_start_loc_long[concept_seq] + local_concept_ids * self.chunk_size - concept_position_ids = position_ids[concept_token_start] + # DCP/Megatron V21 builds HLM rotary embeddings on the compressed + # concept timeline, not on the original token timeline. For ordinary + # full-prompt prefill this yields 0, 1, 2, ...; deriving it from token + # position ids keeps non-zero-offset chunks on the same absolute + # concept index. + concept_position_ids = torch.div( + position_ids[concept_token_start].to(dtype=torch.long), + self.chunk_size, + rounding_mode='floor', + ) max_concepts_per_request = self._get_max_concepts_per_request(token_attn_metadata, concept_q_seqlens_long) return _PrefillConceptLayout( diff --git a/tests/pytorch/backends/test_conceptlm_runtime_ops.py b/tests/pytorch/backends/test_conceptlm_runtime_ops.py new file mode 100644 index 0000000000..32358a7aca --- /dev/null +++ b/tests/pytorch/backends/test_conceptlm_runtime_ops.py @@ -0,0 +1,60 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from types import SimpleNamespace + +import torch + +from lmdeploy.pytorch.backends.default.conceptlm import DefaultConceptLMRuntimeOpsImpl + + +def _make_ops() -> DefaultConceptLMRuntimeOpsImpl: + config = SimpleNamespace( + concept_chunk_size=4, + concept_chunk_merge_method='meanpooling', + concept_shift_feature=True, + ) + return DefaultConceptLMRuntimeOpsImpl(config) + + +def _make_prefill_attn_metadata(q_seqlens: list[int]): + q_seqlens_tensor = torch.tensor(q_seqlens, dtype=torch.int32) + cu_seqlens = torch.nn.functional.pad(torch.cumsum(q_seqlens_tensor, dim=0, dtype=torch.int32), (1, 0)) + return SimpleNamespace( + is_decoding=False, + q_seqlens=q_seqlens_tensor, + q_start_loc=cu_seqlens[:-1], + cu_seqlens_q=cu_seqlens, + kv_seqlens=q_seqlens_tensor, + max_q_seqlen=max(q_seqlens), + ) + + +def test_concept_prefill_position_ids_use_compressed_timeline(): + """HLM RoPE uses concept indices, not token-start positions.""" + ops = _make_ops() + metadata = _make_prefill_attn_metadata([9, 10]) + position_ids = torch.tensor(list(range(9)) + list(range(10)), dtype=torch.long) + + concept_metadata = ops.build_prefill_metadata(metadata, position_ids) + + assert concept_metadata.concept_q_seqlens.tolist() == [2, 2] + assert concept_metadata.concept_position_ids.tolist() == [0, 1, 0, 1] + + +def test_concept_prefill_position_ids_preserve_absolute_prefill_offset(): + """Chunked prefill keeps absolute concept index within the request.""" + ops = _make_ops() + metadata = _make_prefill_attn_metadata([9]) + position_ids = torch.arange(8, 17, dtype=torch.long) + + concept_metadata = ops.build_prefill_metadata(metadata, position_ids) + + assert concept_metadata.concept_q_seqlens.tolist() == [2] + assert concept_metadata.concept_position_ids.tolist() == [2, 3] + + +def test_concept_decode_position_ids_use_current_concept_index(): + ops = _make_ops() + + position_ids = torch.tensor([3, 4, 7, 8], dtype=torch.long) + + assert ops.decode_concept_position_ids(position_ids).tolist() == [0, 0, 1, 1] From 958662912969f8c642cc7534f706c913f690965a Mon Sep 17 00:00:00 2001 From: grimoire Date: Wed, 29 Jul 2026 12:25:51 +0800 Subject: [PATCH 13/16] Fix ConceptLM decode metadata max length --- .../pytorch/backends/default/conceptlm.py | 13 +++- .../backends/test_conceptlm_runtime_ops.py | 60 ------------------- 2 files changed, 12 insertions(+), 61 deletions(-) delete mode 100644 tests/pytorch/backends/test_conceptlm_runtime_ops.py diff --git a/lmdeploy/pytorch/backends/default/conceptlm.py b/lmdeploy/pytorch/backends/default/conceptlm.py index 2253c2d8c8..a6e8d45e22 100644 --- a/lmdeploy/pytorch/backends/default/conceptlm.py +++ b/lmdeploy/pytorch/backends/default/conceptlm.py @@ -196,7 +196,18 @@ def build_concept_decode_metadata_static(self, token_attn_metadata: Any, if hasattr(token_attn_metadata, 'max_q_seqlen'): updates['max_q_seqlen'] = 1 if hasattr(token_attn_metadata, 'max_kv_seqlen'): - updates['max_kv_seqlen'] = getattr(token_attn_metadata, 'max_kv_seqlen') + max_token_kv_seqlen = getattr(token_attn_metadata, 'max_kv_seqlen') + if max_token_kv_seqlen is None: + max_concept_kv_seqlen = ( + int(concept_kv_seqlens.max().item()) if concept_kv_seqlens.numel() > 0 else 1 + ) + else: + max_concept_kv_seqlen = self.concept_count_from_seq_len( + int(max_token_kv_seqlen), + self.chunk_size, + ) + max_concept_kv_seqlen = max(max_concept_kv_seqlen, 1) + updates['max_kv_seqlen'] = max_concept_kv_seqlen if hasattr(token_attn_metadata, 'scheduler_metadata'): updates['scheduler_metadata'] = None if hasattr(token_attn_metadata, 'tile_scheduler_metadata'): diff --git a/tests/pytorch/backends/test_conceptlm_runtime_ops.py b/tests/pytorch/backends/test_conceptlm_runtime_ops.py deleted file mode 100644 index 32358a7aca..0000000000 --- a/tests/pytorch/backends/test_conceptlm_runtime_ops.py +++ /dev/null @@ -1,60 +0,0 @@ -# Copyright (c) OpenMMLab. All rights reserved. -from types import SimpleNamespace - -import torch - -from lmdeploy.pytorch.backends.default.conceptlm import DefaultConceptLMRuntimeOpsImpl - - -def _make_ops() -> DefaultConceptLMRuntimeOpsImpl: - config = SimpleNamespace( - concept_chunk_size=4, - concept_chunk_merge_method='meanpooling', - concept_shift_feature=True, - ) - return DefaultConceptLMRuntimeOpsImpl(config) - - -def _make_prefill_attn_metadata(q_seqlens: list[int]): - q_seqlens_tensor = torch.tensor(q_seqlens, dtype=torch.int32) - cu_seqlens = torch.nn.functional.pad(torch.cumsum(q_seqlens_tensor, dim=0, dtype=torch.int32), (1, 0)) - return SimpleNamespace( - is_decoding=False, - q_seqlens=q_seqlens_tensor, - q_start_loc=cu_seqlens[:-1], - cu_seqlens_q=cu_seqlens, - kv_seqlens=q_seqlens_tensor, - max_q_seqlen=max(q_seqlens), - ) - - -def test_concept_prefill_position_ids_use_compressed_timeline(): - """HLM RoPE uses concept indices, not token-start positions.""" - ops = _make_ops() - metadata = _make_prefill_attn_metadata([9, 10]) - position_ids = torch.tensor(list(range(9)) + list(range(10)), dtype=torch.long) - - concept_metadata = ops.build_prefill_metadata(metadata, position_ids) - - assert concept_metadata.concept_q_seqlens.tolist() == [2, 2] - assert concept_metadata.concept_position_ids.tolist() == [0, 1, 0, 1] - - -def test_concept_prefill_position_ids_preserve_absolute_prefill_offset(): - """Chunked prefill keeps absolute concept index within the request.""" - ops = _make_ops() - metadata = _make_prefill_attn_metadata([9]) - position_ids = torch.arange(8, 17, dtype=torch.long) - - concept_metadata = ops.build_prefill_metadata(metadata, position_ids) - - assert concept_metadata.concept_q_seqlens.tolist() == [2] - assert concept_metadata.concept_position_ids.tolist() == [2, 3] - - -def test_concept_decode_position_ids_use_current_concept_index(): - ops = _make_ops() - - position_ids = torch.tensor([3, 4, 7, 8], dtype=torch.long) - - assert ops.decode_concept_position_ids(position_ids).tolist() == [0, 0, 1, 1] From d3a0a1f20def903d0813212b385704c5dfcca772 Mon Sep 17 00:00:00 2001 From: grimoire Date: Wed, 12 Aug 2026 16:56:59 +0800 Subject: [PATCH 14/16] fix oob --- lmdeploy/pytorch/backends/default/conceptlm.py | 12 ++++++++++++ 1 file changed, 12 insertions(+) diff --git a/lmdeploy/pytorch/backends/default/conceptlm.py b/lmdeploy/pytorch/backends/default/conceptlm.py index a6e8d45e22..92ab2b9838 100644 --- a/lmdeploy/pytorch/backends/default/conceptlm.py +++ b/lmdeploy/pytorch/backends/default/conceptlm.py @@ -179,6 +179,12 @@ def build_concept_decode_metadata_static(self, token_attn_metadata: Any, self.chunk_size, rounding_mode='floor', ).clamp(min=1).to(dtype=kv_dtype) + # CUDAGraph pads decode batches with state_id=-1 rows. Those rows + # still execute concept attention and KV snapshot/restore, so keep + # their concept timeline inside the dummy cache slot. + valid_state_mask = decode_metadata.valid_state_mask.to(device=device) + safe_concept_kv_seqlens = torch.ones_like(concept_kv_seqlens) + concept_kv_seqlens = torch.where(valid_state_mask, concept_kv_seqlens, safe_concept_kv_seqlens) updates = dict( is_decoding=True, @@ -253,6 +259,12 @@ def build_concept_chunk_input( ) concept_attn_metadata = self.build_concept_decode_metadata_static(token_attn_metadata, decode_metadata) concept_position_ids = self.decode_concept_position_ids(decode_metadata.position_ids) + # Match the safe dummy KV length above for graph-padded rows. + concept_position_ids = torch.where( + decode_metadata.valid_state_mask, + concept_position_ids, + torch.zeros_like(concept_position_ids), + ) return ConceptChunkInput( source_states=concept_states, position_ids=concept_position_ids, From 2082551d0b5bd8345be8436301f0402b4ea7220e Mon Sep 17 00:00:00 2001 From: grimoire Date: Sat, 29 Aug 2026 13:51:47 +0800 Subject: [PATCH 15/16] Support sparse ConceptLM HF exports --- lmdeploy/pytorch/configurations/conceptlm.py | 182 ++++++++++++++++++ lmdeploy/pytorch/models/intern_ncp/modules.py | 21 +- tests/pytorch/config/test_model_config.py | 113 +++++++++++ 3 files changed, 299 insertions(+), 17 deletions(-) diff --git a/lmdeploy/pytorch/configurations/conceptlm.py b/lmdeploy/pytorch/configurations/conceptlm.py index 4ba7d64e5b..a38b153007 100644 --- a/lmdeploy/pytorch/configurations/conceptlm.py +++ b/lmdeploy/pytorch/configurations/conceptlm.py @@ -1,4 +1,6 @@ # Copyright (c) OpenMMLab. All rights reserved. +from pathlib import Path + import torch from lmdeploy.pytorch.config import StateCacheSpec @@ -16,6 +18,76 @@ 'concept_last_state', ) +_TRAINING_CONFIG_NAME = 'training_config.yaml' +_CONCEPT_TRAINING_CONFIG_KEYS = { + 'concept_chunk_size': 'conceptlm_chunk_size', + 'concept_shift_feature': 'conceptlm_shift_feature', + 'concept_chunk_merge_method': 'conceptlm_chunk_merge_method', + 'concept_layer_norm_option': 'conceptlm_layer_norm_option', + 'concept_fusion_norm_alpha_init': 'conceptlm_fusion_norm_alpha_init', + 'concept_fusion_alpha_init': 'conceptlm_fusion_alpha_init', + 'concept_hlm_ffn_hidden_size': 'conceptlm_hlm_ffn_hidden_size', + 'concept_hlm_attention_mode': 'conceptlm_hlm_attention_mode', + 'concept_v22_vq_codebook_size': 'conceptlm_v22_vq_codebook_size', + 'concept_v22_vq_num_codebooks': 'conceptlm_v22_vq_num_codebooks', + 'concept_v22_vq_commitment_cost': 'conceptlm_v22_vq_commitment_cost', + 'concept_v22_vq_merge_mode': 'conceptlm_v22_vq_merge_mode', + 'concept_v22_vq_hlm_loss_type': 'conceptlm_v22_vq_hlm_loss_type', + 'concept_v22_vq_detach_hlm_target': 'conceptlm_v22_vq_detach_hlm_target', + 'concept_dd_two_route_add': 'conceptlm_v21_dd_two_route_add', + 'concept_dd_two_route_add_concept_source': 'conceptlm_v21_dd_two_route_add_concept_source', + 'concept_dd_two_route_add_enable_raw_concept_route': ( + 'conceptlm_v21_dd_two_route_add_enable_raw_concept_route' + ), + 'concept_dd_two_route_add_enable_final_concept_route': ( + 'conceptlm_v21_dd_two_route_add_enable_final_concept_route' + ), + 'concept_dd_two_route_add_beta_init': 'conceptlm_v21_dd_two_route_add_beta_init', + 'concept_dd_two_route_add_every_n_layers': 'conceptlm_v21_dd_two_route_add_every_n_layers', + 'concept_dd_two_route_add_concept_route_first_n': 'conceptlm_v21_dd_two_route_add_concept_route_first_n', + 'concept_dd_two_route_add_decoder_hidden_size': 'conceptlm_v21_dd_two_route_add_decoder_hidden_size', + 'concept_dd_two_route_add_concept_hidden_size': 'conceptlm_v21_dd_two_route_add_concept_hidden_size', + 'concept_dd_two_route_add_use_softmax': 'conceptlm_v21_dd_two_route_add_use_softmax', + 'concept_dd_two_route_add_disable_decoder_dd': 'conceptlm_v21_dd_two_route_add_disable_decoder_dd', + 'concept_dd_two_route_add_decoder_use_layernorm': 'conceptlm_v21_dd_two_route_add_decoder_use_layernorm', + 'concept_dd_two_route_add_decoder_use_softmax': 'conceptlm_v21_dd_two_route_add_decoder_use_softmax', + 'concept_dd_two_route_add_concept_use_layernorm': 'conceptlm_v21_dd_two_route_add_concept_use_layernorm', + 'concept_dd_encoder_self_dd': 'conceptlm_v21_dd_encoder_self_dd', + 'concept_dd_encoder_self_dd_every_n_layers': 'conceptlm_v21_dd_encoder_self_dd_every_n_layers', + 'concept_dd_encoder_self_dd_hidden_size': 'conceptlm_v21_dd_encoder_self_dd_hidden_size', + 'concept_dd_encoder_self_dd_use_layernorm': 'conceptlm_v21_dd_encoder_self_dd_use_layernorm', + 'concept_dd_concept_self_dd': 'conceptlm_v21_dd_concept_self_dd', + 'concept_dd_concept_self_dd_every_n_layers': 'conceptlm_v21_dd_concept_self_dd_every_n_layers', + 'concept_dd_concept_self_dd_hidden_size': 'conceptlm_v21_dd_concept_self_dd_hidden_size', + 'concept_dd_concept_self_dd_use_layernorm': 'conceptlm_v21_dd_concept_self_dd_use_layernorm', + 'concept_enable_concept_read_encoder': 'conceptlm_v21_enable_concept_read_encoder', + 'concept_enable_decoder_read_encoder': 'conceptlm_v21_enable_decoder_read_encoder', + 'concept_enable_decoder_read_concept': 'conceptlm_v21_enable_decoder_read_concept', + 'concept_read_encoder_first_n': 'conceptlm_v21_concept_read_encoder_first_n', + 'concept_decoder_read_encoder_first_n': 'conceptlm_v21_decoder_read_encoder_first_n', + 'concept_residual_flow_beta_init': 'conceptlm_v21_residual_flow_beta_init', + 'concept_residual_flow_route_hidden_size': 'conceptlm_v21_residual_flow_route_hidden_size', + 'concept_residual_flow_route_use_softmax': 'conceptlm_v21_residual_flow_route_use_softmax', + 'concept_residual_flow_source_use_layernorm': 'conceptlm_v21_residual_flow_source_use_layernorm', + 'concept_residual_flow_shared_source_norm': 'conceptlm_v21_residual_flow_shared_source_norm', + 'concept_final_read_concept_gate': 'conceptlm_v21_final_read_concept_gate', + 'concept_final_read_concept_gate_init_final': 'conceptlm_v21_final_read_concept_gate_init_final', + 'concept_final_read_concept_gate_target_final': 'conceptlm_v21_final_read_concept_gate_target_final', + 'concept_final_read_concept_gate_reg_weight': 'conceptlm_v21_final_read_concept_gate_reg_weight', + 'concept_dd_self_dd_mode': 'conceptlm_v21_dd_self_dd_mode', + 'concept_enable_full_residual_flow': 'conceptlm_v21_enable_full_residual_flow', + 'concept_encoder_layers': 'conceptlm_encoder_layers', + 'concept_special_layers': 'conceptlm_special_layers', + 'concept_decoder_layers': 'conceptlm_decoder_layers', +} +_REQUIRED_CONCEPT_CONFIG_KEYS = ( + 'concept_encoder_layers', + 'concept_special_layers', + 'concept_decoder_layers', + 'concept_v22_vq_num_codebooks', + 'concept_v22_vq_codebook_size', +) + def _get_concept_state_dtype(hf_config): """Return the dtype used by ConceptLM sequence-state caches.""" @@ -30,6 +102,112 @@ def _get_concept_state_dtype(hf_config): return torch.float16 +def _training_config_value(value): + """Unwrap wandb-style ``{desc, value}`` entries.""" + if isinstance(value, dict) and 'value' in value: + return value['value'] + return value + + +def _load_training_config(model_path: str = None): + """Load ConceptLM's exported training config when the HF config is + sparse.""" + if not model_path: + return None + + path = Path(model_path) / _TRAINING_CONFIG_NAME + if not path.is_file(): + return None + + try: + import yaml + with path.open(encoding='utf-8') as f: + data = yaml.safe_load(f) + except Exception as e: # noqa: BLE001 + logger.warning(f'ConceptLM: failed to load {path}: {e}') + return None + + return data if isinstance(data, dict) else None + + +def _fill_concept_runtime_config(hf_config, model_path: str = None): + """Fill old LMDeploy ConceptLM fields from the new exported training + config. + + Newer ConceptLM HF exports keep only generic model metadata in + ``config.json`` and leave ConceptLM runtime structure in + ``training_config.yaml`` using the original training argument names. The + LMDeploy model uses the older ``concept_*`` names internally, so normalize + the config once at the configuration boundary. + """ + training_config = _load_training_config(model_path) + if training_config is None: + return + + for attr_name, training_name in _CONCEPT_TRAINING_CONFIG_KEYS.items(): + if hasattr(hf_config, attr_name) or training_name not in training_config: + continue + setattr(hf_config, attr_name, _training_config_value(training_config[training_name])) + + +def _check_concept_runtime_config(hf_config): + """Raise a readable error for sparse ConceptLM exports.""" + missing = [name for name in _REQUIRED_CONCEPT_CONFIG_KEYS if not hasattr(hf_config, name)] + if missing: + raise AttributeError( + 'ConceptLM config is missing runtime fields ' + f'{missing}. If this is a new HF export, keep {_TRAINING_CONFIG_NAME} beside config.json ' + 'so LMDeploy can backfill the ConceptLM architecture fields.') + + +def _normalize_concept_rotary_config(hf_config): + """Normalize ConceptLM rotary fields for the shared rotary builder.""" + head_dim = getattr(hf_config, 'head_dim', None) + if head_dim is None: + head_dim = getattr(hf_config, 'kv_channels', None) + if head_dim is None: + head_dim = hf_config.hidden_size // hf_config.num_attention_heads + hf_config.head_dim = int(head_dim) + + if not hasattr(hf_config, 'rope_theta') and hasattr(hf_config, 'rotary_base'): + hf_config.rope_theta = hf_config.rotary_base + + rotary_percent = getattr(hf_config, 'rotary_percent', None) + if rotary_percent is not None and not hasattr(hf_config, 'partial_rotary_factor'): + rotary_dim = int(hf_config.head_dim * float(rotary_percent)) + rotary_dim -= rotary_dim % 2 + if rotary_dim <= 0: + raise ValueError( + f'Invalid ConceptLM rotary dimension: head_dim={hf_config.head_dim}, rotary_percent={rotary_percent}') + hf_config.partial_rotary_factor = rotary_dim / hf_config.head_dim + + position_embedding_type = getattr(hf_config, 'position_embedding_type', 'rope') + if position_embedding_type != 'yarn': + return + + scaling_factor = getattr(hf_config, 'yarn_rotary_scaling_factor', 1.0) + if scaling_factor is None: + scaling_factor = 1.0 + rope_scaling = { + 'rope_type': 'yarn', + 'rope_theta': getattr(hf_config, 'rope_theta', getattr(hf_config, 'rotary_base', 10000)), + 'factor': float(scaling_factor), + 'beta_fast': getattr(hf_config, 'yarn_beta_fast', 32.0), + 'beta_slow': getattr(hf_config, 'yarn_beta_slow', 1.0), + 'mscale': getattr(hf_config, 'yarn_mscale', 1.0), + 'mscale_all_dim': getattr(hf_config, 'yarn_mscale_all_dim', 0.0), + 'truncate': getattr(hf_config, 'yarn_correction_range_round_to_int', True), + } + original_max_position_embeddings = getattr(hf_config, 'yarn_original_max_position_embeddings', + getattr(hf_config, 'max_position_embeddings', None)) + if original_max_position_embeddings is not None: + rope_scaling['original_max_position_embeddings'] = original_max_position_embeddings + if getattr(hf_config, 'rope_parameters', None) is None: + hf_config.rope_parameters = dict(rope_scaling) + if getattr(hf_config, 'rope_scaling', None) is None: + hf_config.rope_scaling = dict(rope_scaling) + + class ConceptLMModelConfigBuilder(AutoModelConfigBuilder): """Config builder for ConceptLM V2.2-VQ. @@ -46,6 +224,10 @@ def condition(cls, hf_config): @classmethod def build(cls, hf_config, model_path: str = None, **kwargs): """build.""" + _fill_concept_runtime_config(hf_config, model_path) + _check_concept_runtime_config(hf_config) + _normalize_concept_rotary_config(hf_config) + # fill missing special token ids from the tokenizer if getattr(hf_config, 'bos_token_id', None) is None: cls._fill_special_tokens(hf_config, model_path) diff --git a/lmdeploy/pytorch/models/intern_ncp/modules.py b/lmdeploy/pytorch/models/intern_ncp/modules.py index da16de9d31..bb06b9b5fa 100644 --- a/lmdeploy/pytorch/models/intern_ncp/modules.py +++ b/lmdeploy/pytorch/models/intern_ncp/modules.py @@ -12,7 +12,7 @@ Attention, RMSNorm, SiluAndMul, - build_rotary_embedding, + build_rotary_embedding_from_config, ) from lmdeploy.pytorch.nn.linear import ( build_down_linear, @@ -31,6 +31,7 @@ _HistoryStates = torch.Tensor _SourceStates = torch.Tensor | None + def _get_configured_window(config: PretrainedConfig): """Return the reference OLMo window setting as an int or None.""" window_size = getattr(config, 'window_size', None) @@ -51,22 +52,8 @@ def _make_olmo_rotary_embedding(config: PretrainedConfig, if rotary_interleaved: raise NotImplementedError('ConceptLM rotary_interleaved=True is not supported by the LMDeploy block yet.') - head_dim = int(config.kv_channels) - rotary_percent = float(getattr(config, 'rotary_percent', 1.0)) - rotary_dim = int(head_dim * rotary_percent) - rotary_dim -= rotary_dim % 2 - if rotary_dim <= 0: - raise ValueError(f'Invalid ConceptLM rotary dimension: head_dim={head_dim}, rotary_percent={rotary_percent}') - - partial_rotary_factor = rotary_dim / head_dim - return build_rotary_embedding( - dim=head_dim, - max_position_embeddings=getattr(config, 'max_position_embeddings', getattr(config, 'max_sequence_length', - 2048)), - base=getattr(config, 'rotary_base', 10000), - partial_rotary_factor=partial_rotary_factor, - device=device, - ) + return build_rotary_embedding_from_config(config, device=device) + def _qk_rmsnorm_variance(query: torch.Tensor, key: torch.Tensor) -> torch.Tensor: """Local Q/K squared sums before TP all-reduce.""" diff --git a/tests/pytorch/config/test_model_config.py b/tests/pytorch/config/test_model_config.py index c4c94edb42..5db039a8e0 100644 --- a/tests/pytorch/config/test_model_config.py +++ b/tests/pytorch/config/test_model_config.py @@ -1,3 +1,4 @@ +from textwrap import dedent from types import SimpleNamespace import pytest @@ -6,6 +7,7 @@ from lmdeploy.pytorch.config import CacheConfig, DistConfig, ModelConfig from lmdeploy.pytorch.configurations import AutoModelConfigBuilder from lmdeploy.pytorch.configurations.deepseek_v4 import update_cache_config as update_deepseek_v4_cache_config +from lmdeploy.pytorch.nn import RopeType, build_rotary_embedding_from_config, build_rotary_params def _make_model_config(num_attention_heads=32, num_key_value_heads=8, dist_config=None): @@ -37,6 +39,70 @@ def _make_deepseek_v4_hf_config(compress_ratios, num_hidden_layers=3): ) +def _make_sparse_conceptlm_hf_config(): + return SimpleNamespace( + architectures=['ConceptLMV22VQForCausalLM'], + model_type='conceptlm_v22_vq', + hidden_size=4096, + num_hidden_layers=32, + num_layers=32, + num_attention_heads=32, + kv_channels=128, + bos_token_id=100257, + eos_token_id=100257, + pad_token_id=100277, + vocab_size=100278, + torch_dtype='bfloat16', + concept_chunk_size=4, + concept_shift_feature=True, + position_embedding_type='yarn', + max_position_embeddings=65536, + max_sequence_length=65536, + rotary_base=500000, + rotary_percent=1.0, + yarn_rotary_scaling_factor=8.0, + yarn_original_max_position_embeddings=8192, + yarn_beta_fast=32.0, + yarn_beta_slow=1.0, + yarn_mscale=1.0, + yarn_mscale_all_dim=0.0, + yarn_correction_range_round_to_int=True, + ) + + +def _write_conceptlm_training_config(model_dir): + (model_dir / 'training_config.yaml').write_text( + dedent(""" + conceptlm_encoder_layers: + value: 16 + conceptlm_special_layers: + value: 8 + conceptlm_decoder_layers: + value: 16 + conceptlm_chunk_merge_method: + value: meanpooling + conceptlm_fusion_norm_alpha_init: + value: 0.1 + conceptlm_v22_vq_codebook_size: + value: 128 + conceptlm_v22_vq_num_codebooks: + value: 32 + conceptlm_v21_dd_two_route_add_decoder_use_softmax: + value: true + conceptlm_v21_enable_concept_read_encoder: + value: true + conceptlm_v21_enable_decoder_read_encoder: + value: true + conceptlm_v21_enable_decoder_read_concept: + value: true + conceptlm_v21_concept_read_encoder_first_n: + value: -1 + conceptlm_v21_decoder_read_encoder_first_n: + value: -1 + """), + encoding='utf-8') + + def test_get_num_qkv_head_by_tp_from_dist_config(): model_config = _make_model_config(dist_config=DistConfig(tp=4)) @@ -68,6 +134,53 @@ def test_from_hf_config_keeps_dist_config_for_head_split(): assert model_config.get_num_qkv_head_by_tp() == (8, 2) +def test_conceptlm_model_config_backfills_new_export_training_config(tmp_path): + hf_config = _make_sparse_conceptlm_hf_config() + _write_conceptlm_training_config(tmp_path) + + model_config = AutoModelConfigBuilder.build(hf_config, str(tmp_path)) + + assert hf_config.concept_encoder_layers == 16 + assert hf_config.concept_special_layers == 8 + assert hf_config.concept_decoder_layers == 16 + assert hf_config.concept_v22_vq_codebook_size == 128 + assert hf_config.concept_v22_vq_num_codebooks == 32 + assert hf_config.concept_chunk_merge_method == 'meanpooling' + assert model_config.num_layers == 40 + assert model_config.states_shapes == [((16, 4096), torch.float32), ((9, 4096), torch.bfloat16)] + + +def test_conceptlm_model_config_reports_sparse_export_without_training_config(tmp_path): + hf_config = _make_sparse_conceptlm_hf_config() + + with pytest.raises(AttributeError, match='training_config.yaml'): + AutoModelConfigBuilder.build(hf_config, str(tmp_path)) + + +def test_conceptlm_yarn_rotary_config_uses_original_context_length(tmp_path): + hf_config = _make_sparse_conceptlm_hf_config() + _write_conceptlm_training_config(tmp_path) + model_config = AutoModelConfigBuilder.build(hf_config, str(tmp_path)) + + rotary_params = build_rotary_params(hf_config) + + assert model_config.head_dim == 128 + assert hf_config.head_dim == 128 + assert hf_config.rope_theta == 500000 + assert hf_config.partial_rotary_factor == 1.0 + assert hf_config.rope_scaling['rope_type'] == 'yarn' + assert hf_config.rope_parameters['rope_theta'] == 500000 + assert rotary_params['emb_type'] is RopeType.Yarn + assert rotary_params['scaling_factor'] == 8.0 + assert rotary_params['max_position_embeddings'] == 8192 + assert rotary_params['yarn_params'].beta_fast == 32.0 + assert rotary_params['yarn_params'].beta_slow == 1.0 + assert rotary_params['yarn_params'].mscale == 1.0 + assert rotary_params['yarn_params'].mscale_all_dim == 0.0 + assert rotary_params['yarn_params'].truncate is True + assert build_rotary_embedding_from_config(hf_config).base == 500000 + + def test_get_num_qkv_head_by_tp_with_dist_config_tp(): model_config = _make_model_config(dist_config=DistConfig(tp=2)) From 5976009994a7d40a74eaebe4fce401b6382d0922 Mon Sep 17 00:00:00 2001 From: grimoire Date: Fri, 4 Sep 2026 15:47:37 +0800 Subject: [PATCH 16/16] Address ConceptLM review comments --- lmdeploy/pytorch/config.py | 3 +++ lmdeploy/pytorch/configurations/conceptlm.py | 7 +++--- lmdeploy/pytorch/models/intern_ncp/modules.py | 3 +-- tests/pytorch/config/test_model_config.py | 22 +++++++++++++++++++ 4 files changed, 30 insertions(+), 5 deletions(-) diff --git a/lmdeploy/pytorch/config.py b/lmdeploy/pytorch/config.py index 6125b2a532..527a5fa75c 100644 --- a/lmdeploy/pytorch/config.py +++ b/lmdeploy/pytorch/config.py @@ -540,6 +540,7 @@ def from_pretrained( model_config = cls.from_hf_config( hf_config, pretrained_model_name_or_path, + trust_remote_code=trust_remote_code, dtype=dtype, dist_config=dist_config, is_draft_model=is_draft_model, @@ -576,6 +577,7 @@ def from_hf_config( spec_method: str = None, device_type: str = 'auto', num_spec_tokens: int = 0, + trust_remote_code: bool = False, ): """From huggingface config.""" from lmdeploy.pytorch.configurations import AutoModelConfigBuilder @@ -591,6 +593,7 @@ def from_hf_config( spec_method=spec_method, num_spec_tokens=num_spec_tokens, device_type=device_type, + trust_remote_code=trust_remote_code, ) if model_config.k_head_dim is None: diff --git a/lmdeploy/pytorch/configurations/conceptlm.py b/lmdeploy/pytorch/configurations/conceptlm.py index a38b153007..74404c4895 100644 --- a/lmdeploy/pytorch/configurations/conceptlm.py +++ b/lmdeploy/pytorch/configurations/conceptlm.py @@ -224,13 +224,14 @@ def condition(cls, hf_config): @classmethod def build(cls, hf_config, model_path: str = None, **kwargs): """build.""" + trust_remote_code = bool(kwargs.get('trust_remote_code', False)) _fill_concept_runtime_config(hf_config, model_path) _check_concept_runtime_config(hf_config) _normalize_concept_rotary_config(hf_config) # fill missing special token ids from the tokenizer if getattr(hf_config, 'bos_token_id', None) is None: - cls._fill_special_tokens(hf_config, model_path) + cls._fill_special_tokens(hf_config, model_path, trust_remote_code=trust_remote_code) model_config = DefaultModelConfigBuilder.build(hf_config, model_path, **kwargs) @@ -284,10 +285,10 @@ def build(cls, hf_config, model_path: str = None, **kwargs): return model_config @staticmethod - def _fill_special_tokens(hf_config, model_path: str = None): + def _fill_special_tokens(hf_config, model_path: str = None, trust_remote_code: bool = False): try: from transformers import AutoTokenizer - tok = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True) + tok = AutoTokenizer.from_pretrained(model_path, trust_remote_code=trust_remote_code) hf_config.bos_token_id = tok.bos_token_id hf_config.eos_token_id = tok.eos_token_id if getattr(hf_config, 'pad_token_id', None) is None: diff --git a/lmdeploy/pytorch/models/intern_ncp/modules.py b/lmdeploy/pytorch/models/intern_ncp/modules.py index bb06b9b5fa..460805a085 100644 --- a/lmdeploy/pytorch/models/intern_ncp/modules.py +++ b/lmdeploy/pytorch/models/intern_ncp/modules.py @@ -675,9 +675,8 @@ def forward(self, past_key_value: list[torch.Tensor] | None = None, attn_metadata: Any = None): """Rewrite of _OlmoSelfAttention.forward.""" - # qkv proj -> (batch, seq, num_heads, head_dim) each qkv_states = self.qkv_proj(hidden_states) - qkv_states = qkv_states.flatten(0, -2) # (-1, heads_total, head_dim) + qkv_states = qkv_states.flatten(0, -2) # [num_tokens, packed_qkv_dim] query_states, key_states, value_states = self.qkv_proj.split_qkv(qkv_states) # q, k norm (whole hidden_size) diff --git a/tests/pytorch/config/test_model_config.py b/tests/pytorch/config/test_model_config.py index 6db62c058c..ef7ba257bb 100644 --- a/tests/pytorch/config/test_model_config.py +++ b/tests/pytorch/config/test_model_config.py @@ -168,6 +168,28 @@ def test_conceptlm_model_config_reports_sparse_export_without_training_config(tm AutoModelConfigBuilder.build(hf_config, str(tmp_path)) +@pytest.mark.parametrize('trust_remote_code', [False, True]) +def test_conceptlm_special_token_fallback_respects_trust_remote_code(tmp_path, monkeypatch, trust_remote_code): + hf_config = _make_sparse_conceptlm_hf_config() + hf_config.bos_token_id = None + hf_config.eos_token_id = None + hf_config.pad_token_id = None + _write_conceptlm_training_config(tmp_path) + tokenizer_calls = [] + + def fake_from_pretrained(model_path, trust_remote_code=False): + tokenizer_calls.append((model_path, trust_remote_code)) + return SimpleNamespace(bos_token_id=1, eos_token_id=2, pad_token_id=0) + + monkeypatch.setattr('transformers.AutoTokenizer.from_pretrained', fake_from_pretrained) + + model_config = ModelConfig.from_hf_config(hf_config, str(tmp_path), trust_remote_code=trust_remote_code) + + assert tokenizer_calls == [(str(tmp_path), trust_remote_code)] + assert model_config.bos_token_id == 1 + assert model_config.eos_token_id == [2] + + def test_conceptlm_yarn_rotary_config_uses_original_context_length(tmp_path): hf_config = _make_sparse_conceptlm_hf_config() _write_conceptlm_training_config(tmp_path)