diff --git a/lmdeploy/pytorch/backends/conceptlm.py b/lmdeploy/pytorch/backends/conceptlm.py new file mode 100644 index 0000000000..6bd9a4b300 --- /dev/null +++ b/lmdeploy/pytorch/backends/conceptlm.py @@ -0,0 +1,163 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from abc import ABC, abstractmethod +from dataclasses import dataclass +from typing import Any + +import torch +from torch import Tensor +from transformers.configuration_utils import PretrainedConfig + +from .base import BuildSpec + + +@dataclass +class ConceptChunkInput: + """Concept-stream input rows prepared from token-stream encoder states. + + ``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 +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_start_ids: 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 + + +class ConceptLMRuntimeOpsImpl(ABC): + """Backend contract for ConceptLM runtime/cache operations.""" + + 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)) + + @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.""" + raise NotImplementedError('Not implemented.') + + @abstractmethod + 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.""" + raise NotImplementedError('Not implemented.') + + @abstractmethod + 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 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.""" + raise NotImplementedError('Not implemented.') + + @abstractmethod + 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 for prefill or decode.""" + raise NotImplementedError('Not implemented.') + + +@dataclass(frozen=True) +class ConceptLMRuntimeOpsBuildSpec(BuildSpec[ConceptLMRuntimeOpsImpl]): + """Immutable requirements for constructing ConceptLM runtime ops.""" + + config: PretrainedConfig diff --git a/lmdeploy/pytorch/backends/cuda/attention/fa3.py b/lmdeploy/pytorch/backends/cuda/attention/fa3.py index 63735bdc0f..40f29eff87 100644 --- a/lmdeploy/pytorch/backends/cuda/attention/fa3.py +++ b/lmdeploy/pytorch/backends/cuda/attention/fa3.py @@ -325,6 +325,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/conceptlm.py b/lmdeploy/pytorch/backends/cuda/conceptlm.py new file mode 100644 index 0000000000..c00aa248d3 --- /dev/null +++ b/lmdeploy/pytorch/backends/cuda/conceptlm.py @@ -0,0 +1,154 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from torch import Tensor + +from lmdeploy.pytorch.kernels.cuda.conceptlm import ( + decode_chunk_state_update, + decode_concept_state_update, + decode_kv_cache_restore, + decode_kv_cache_snapshot, + prefill_chunk_state_update, + prefill_state_cache_update, +) + +from ..default.conceptlm import DefaultConceptLMRuntimeOpsImpl + + +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_impl( + 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_impl( + 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, + current_source_states: Tensor, + state_ids: Tensor, + position_ids: Tensor, + chunk_size: int, + merge_method: str, + ) -> tuple[Tensor, Tensor]: + """Update state cache and return fixed-shape concept inputs.""" + return 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 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, + ) diff --git a/lmdeploy/pytorch/backends/cuda/op_backend.py b/lmdeploy/pytorch/backends/cuda/op_backend.py index 5847064776..8ce969caa9 100644 --- a/lmdeploy/pytorch/backends/cuda/op_backend.py +++ b/lmdeploy/pytorch/backends/cuda/op_backend.py @@ -39,6 +39,7 @@ def build_op(cls, spec: BuildSpec[ImplT], *, enable_deterministic: bool = False) from ..blockedf8_modules import LinearBlockedF8BuildSpec from ..causal_conv1d import CausalConv1dBuildSpec from ..compressor import V4CompressorBuildSpec + from ..conceptlm import ConceptLMRuntimeOpsBuildSpec from ..flash_attention import FlashAttentionBuildSpec from ..gated_delta_rule import GatedDeltaMetaBuildSpec, GatedDeltaRuleBuildSpec from ..hc_prepost import HCPrePostBuildSpec @@ -213,6 +214,9 @@ def build_op(cls, spec: BuildSpec[ImplT], *, enable_deterministic: bool = False) logit_softcapping=spec.logit_softcapping, ), ) + if isinstance(spec, ConceptLMRuntimeOpsBuildSpec): + from .conceptlm import TritonConceptLMRuntimeOpsImpl + return cast(ImplT, TritonConceptLMRuntimeOpsImpl(spec.config)) return super().build_op(spec, enable_deterministic=enable_deterministic) @staticmethod diff --git a/lmdeploy/pytorch/backends/default/conceptlm.py b/lmdeploy/pytorch/backends/default/conceptlm.py new file mode 100644 index 0000000000..2f11068c45 --- /dev/null +++ b/lmdeploy/pytorch/backends/default/conceptlm.py @@ -0,0 +1,937 @@ +# 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, + 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.""" + + 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) + # 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, + 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'): + 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'): + 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) + # 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, + 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 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 + 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 + # 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( + 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)), + dtype=torch.float32) + 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.float() * valid_tail.to(dtype=torch.float32).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_impl( + 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_impl( + 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, + 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, + current_source_states: Tensor, + state_ids: Tensor, + position_ids: Tensor, + chunk_size: int, + merge_method: str, + ) -> 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, ( + 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 = 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()}.') + + 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) + 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_rows, previous_rows) + concept_input_states = update_rows + elif merge_method == 'last': + update_rows = current_rows + concept_input_states = current_rows + else: + 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_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]) + if state_id >= 0: + chunk_source_state_cache[state_id].copy_(next_rows[batch_idx]) + return concept_input_states, 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)) diff --git a/lmdeploy/pytorch/backends/default/op_backend.py b/lmdeploy/pytorch/backends/default/op_backend.py index 8e8f5c57ee..e412b91609 100644 --- a/lmdeploy/pytorch/backends/default/op_backend.py +++ b/lmdeploy/pytorch/backends/default/op_backend.py @@ -19,6 +19,7 @@ def build_op(cls, spec: BuildSpec[ImplT], *, enable_deterministic: bool = False) from ..activation import GeluAndMulBuildSpec, SiluAndMulBuildSpec from ..apply_rotary_emb import ApplyRotaryEmbBuildSpec from ..awq_modules import LinearW4A16BuildSpec + from ..conceptlm import ConceptLMRuntimeOpsBuildSpec from ..embedding import EmbeddingBuildSpec from ..linear import LinearBuildSpec from ..moe import SoftmaxTopKBuildSpec @@ -89,6 +90,9 @@ def build_op(cls, spec: BuildSpec[ImplT], *, enable_deterministic: bool = False) if isinstance(spec, LinearBuildSpec): from .linear import DefaultLinearImpl return cast(ImplT, DefaultLinearImpl()) + if isinstance(spec, ConceptLMRuntimeOpsBuildSpec): + from .conceptlm import DefaultConceptLMRuntimeOpsImpl + return cast(ImplT, DefaultConceptLMRuntimeOpsImpl(spec.config)) spec_name = type(spec).__name__ raise RuntimeError(f'Build spec {spec_name} is not supported by {cls.get_name()} backend.') diff --git a/lmdeploy/pytorch/config.py b/lmdeploy/pytorch/config.py index 9e0a43568a..e7518ee306 100644 --- a/lmdeploy/pytorch/config.py +++ b/lmdeploy/pytorch/config.py @@ -542,6 +542,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, @@ -578,6 +579,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 @@ -593,6 +595,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 new file mode 100644 index 0000000000..74404c4895 --- /dev/null +++ b/lmdeploy/pytorch/configurations/conceptlm.py @@ -0,0 +1,297 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from pathlib import Path + +import torch + +from lmdeploy.pytorch.config import StateCacheSpec +from lmdeploy.utils import get_logger + +from .builder import AutoModelConfigBuilder +from .default import DefaultModelConfigBuilder + +logger = get_logger('lmdeploy') + +CONCEPT_STATE_CHUNK_SOURCE = 0 +CONCEPT_STATE_LAST = 1 +CONCEPT_STATE_NAMES = ( + 'concept_chunk_source_state', + '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.""" + 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 + + +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. + + 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.""" + 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, trust_remote_code=trust_remote_code) + + 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 + + hidden_size = int(hf_config.hidden_size) + 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. 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 + # 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), torch.float32), + StateCacheSpec(CONCEPT_STATE_NAMES[CONCEPT_STATE_LAST], (concept_last_state_sources, hidden_size), + last_state_dtype), + ] + 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_idx = CONCEPT_STATE_LAST + return model_config + + @staticmethod + 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=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: + 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/kernels/cuda/conceptlm.py b/lmdeploy/pytorch/kernels/cuda/conceptlm.py new file mode 100644 index 0000000000..0c3b959734 --- /dev/null +++ b/lmdeploy/pytorch/kernels/cuda/conceptlm.py @@ -0,0 +1,725 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""ConceptLM runtime kernels.""" + +import torch +import triton +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, + current_states, + state_ids, + position_ids, + concept_inputs, + 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, + 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) + + concept_ptrs = concept_inputs + batch_id * out_stride_b + source_id * out_stride_s + hidden_id * out_stride_h + tl.store(concept_ptrs, concept_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) + + +@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: + 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 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, + 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, 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.' + 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) + 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, + update_mask, + *chunk_source_state_cache.stride(), + *current_source_states.stride(), + *concept_inputs.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, 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/__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..cd1dc7f045 --- /dev/null +++ b/lmdeploy/pytorch/models/intern_ncp/metadata.py @@ -0,0 +1,191 @@ +# 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) + + +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] diff --git a/lmdeploy/pytorch/models/intern_ncp/modeling.py b/lmdeploy/pytorch/models/intern_ncp/modeling.py new file mode 100644 index 0000000000..3cd07b82ca --- /dev/null +++ b/lmdeploy/pytorch/models/intern_ncp/modeling.py @@ -0,0 +1,603 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from collections.abc import Iterable, Mapping +from dataclasses import dataclass +from typing import Any + +import torch +from torch import nn +from transformers.configuration_utils import PretrainedConfig + +from lmdeploy.pytorch.backends.conceptlm import ( + ConceptChunkInput, + ConceptDecoderInput, + ConceptRuntimeCaches, +) +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, + ConceptMetadata, +) +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 _EncoderOutput: + """Encoder result plus its reusable layer-major SelfDD history buffer.""" + + hidden_states: torch.Tensor + raw_states: list[torch.Tensor] + history_buffer: 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(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 + # ``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 + + hidden_states, position_ids = self._normalize_forward_inputs( + hidden_states, + position_ids, + concept_metadata, + ) + self._validate_concept_caches(concept_metadata, concept_caches) + return self._forward_token_stream( + hidden_states, + 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, + ) + + @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: + """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_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, + 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 = 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 _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_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 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) + + 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=chunk_input.position_ids, + attn_metadata=chunk_input.attn_metadata, + ) + + 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 _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: ConceptDecoderInput, + 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: 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) + + concept_states = self.decoder_read_concept_shared_source_norm(decoder_concepts.route_states) + return decoder_encoder_states, concept_states + + 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, + 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 + 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.end_concept_forward( + chunk_input, + runtime_caches, + forward_context, + source_states, + concept_output.predicted_vectors, + concept_output.raw_states, + ) + decoder_concepts = self.concept_ops.build_decoder_concept_input( + chunk_input, + runtime_caches, + forward_context, + concept_output.predicted_vectors, + concept_output.raw_states, + ) + 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 _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: ConceptDecoderInput, + 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..460805a085 --- /dev/null +++ b/lmdeploy/pytorch/models/intern_ncp/modules.py @@ -0,0 +1,970 @@ +# 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_config, +) +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.') + + 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.""" + 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_states = self.qkv_proj(hidden_states) + 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) + 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 or skip_frequency is None: + return None + if 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/models/module_map.py b/lmdeploy/pytorch/models/module_map.py index 48e54331be..dbc0d154fb 100644 --- a/lmdeploy/pytorch/models/module_map.py +++ b/lmdeploy/pytorch/models/module_map.py @@ -290,6 +290,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', diff --git a/lmdeploy/pytorch/nn/__init__.py b/lmdeploy/pytorch/nn/__init__.py index c167e588ec..c39db60d30 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, ParallelLMHead # 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..b4cb469d9f --- /dev/null +++ b/lmdeploy/pytorch/nn/conceptlm.py @@ -0,0 +1,89 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from typing import Any + +import torch +from torch import Tensor, nn + +from lmdeploy.pytorch.backends import get_backend +from lmdeploy.pytorch.backends.conceptlm import ( + ConceptChunkInput, + ConceptDecoderInput, + ConceptForwardContext, + ConceptLMRuntimeOpsBuildSpec, + ConceptRuntimeCaches, +) + + +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, config): + super().__init__() + self.impl = get_backend().build_op(ConceptLMRuntimeOpsBuildSpec(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_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 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 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.""" + return self.impl.end_concept_forward( + chunk_input, + runtime_caches, + forward_context, + source_states, + predicted_vectors, + concept_raw_states, + ) + + 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 for prefill or decode.""" + return self.impl.build_decoder_concept_input( + chunk_input, + runtime_caches, + forward_context, + predicted_vectors, + concept_raw_states, + ) diff --git a/tests/pytorch/config/test_model_config.py b/tests/pytorch/config/test_model_config.py index f5934da3d8..ef7ba257bb 100644 --- a/tests/pytorch/config/test_model_config.py +++ b/tests/pytorch/config/test_model_config.py @@ -1,12 +1,15 @@ +from textwrap import dedent from types import SimpleNamespace import pytest +import torch from lmdeploy.hf_configs import config_from_pretrained from lmdeploy.hf_configs.configuration_kimi_k2 import KimiK2Config from lmdeploy.pytorch.config import CacheConfig, DistConfig, ModelConfig, QuantizationConfig 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): @@ -47,6 +50,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)) @@ -78,6 +145,75 @@ 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)) + + +@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) + 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 + + @pytest.mark.parametrize( ('num_kv_heads', 'expected_effective_heads', 'expected_replica_num'), [(32, 32, 1), (2, 8, 4), (1, 8, 8)],