Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
20 commits
Select commit Hold shift + click to select a range
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
163 changes: 163 additions & 0 deletions lmdeploy/pytorch/backends/conceptlm.py
Original file line number Diff line number Diff line change
@@ -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
2 changes: 2 additions & 0 deletions lmdeploy/pytorch/backends/cuda/attention/fa3.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
154 changes: 154 additions & 0 deletions lmdeploy/pytorch/backends/cuda/conceptlm.py
Original file line number Diff line number Diff line change
@@ -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,
)
4 changes: 4 additions & 0 deletions lmdeploy/pytorch/backends/cuda/op_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
Loading
Loading