From 870d24e8fbb3fa5bbb63c79aa8b9031bd0c56048 Mon Sep 17 00:00:00 2001 From: bltcn Date: Sun, 7 Dec 2025 00:35:55 +0800 Subject: [PATCH 1/7] Create sync.yml --- .github/workflows/sync.yml | 41 ++++++++++++++++++++++++++++++++++++++ 1 file changed, 41 insertions(+) create mode 100644 .github/workflows/sync.yml diff --git a/.github/workflows/sync.yml b/.github/workflows/sync.yml new file mode 100644 index 0000000000..971c17b348 --- /dev/null +++ b/.github/workflows/sync.yml @@ -0,0 +1,41 @@ +name: Upstream Sync + +permissions: + contents: write + +on: + schedule: + - cron: "0 0 * * *" # 每天 UTC 时间 0点运行一次 (你可以修改这个 cron 表达式) + workflow_dispatch: # 允许你手动点击按钮触发 + +jobs: + sync_latest_from_upstream: + name: Sync latest commits from upstream repo + runs-on: ubuntu-latest + if: ${{ github.event.repository.fork }} # 只有当这是一个 Fork 仓库时才运行 + + steps: + # 第一步:检出你的代码 + - name: Checkout target repo + uses: actions/checkout@v3 + + # 第二步:运行同步 Action + - name: Sync upstream changes + id: sync + uses: aormsby/Fork-Sync-With-Upstream-action@v3.4 + with: + upstream_sync_repo: InternLM/lmdeploy # 【重要】请修改为源仓库的 用户名/仓库名 + upstream_sync_branch: main # 【重要】源仓库的分支名 (main 或 master) + target_sync_branch: main # 你想要同步到的本地分支名 + target_repo_token: ${{ secrets.GITHUB_TOKEN }} # 自动生成的 Token,无需修改 + + # 设置为 true 会在发生冲突时导致 Action 失败,并在测试模式下运行 + test_mode: false + + # 第三步:如果同步失败(通常是因为有冲突),打印提示 + - name: Sync check + if: failure() + run: | + echo "[Error] 由于上游仓库的变更与本地变更冲突,无法自动同步。" + echo "请手动解决冲突。" + exit 1 From a05b815748341884b79c5d6cc7b0ec9fcedd28bb Mon Sep 17 00:00:00 2001 From: bltcn Date: Sun, 7 Dec 2025 00:38:21 +0800 Subject: [PATCH 2/7] Update sync.yml --- .github/workflows/sync.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/sync.yml b/.github/workflows/sync.yml index 971c17b348..16d72e0cfa 100644 --- a/.github/workflows/sync.yml +++ b/.github/workflows/sync.yml @@ -5,7 +5,7 @@ permissions: on: schedule: - - cron: "0 0 * * *" # 每天 UTC 时间 0点运行一次 (你可以修改这个 cron 表达式) + - cron: "0 * * * *" # 每天 UTC 时间 0点运行一次 (你可以修改这个 cron 表达式) workflow_dispatch: # 允许你手动点击按钮触发 jobs: From 7be6a7d1bb0da154ec4eaafeee3f8519be9e6efb Mon Sep 17 00:00:00 2001 From: Li Zhang Date: Mon, 28 Sep 2026 08:55:50 +0000 Subject: [PATCH 3/7] feat(turbomind): EAGLE3 and Qwen3.5 MTP speculative decoding - Fixed-chain speculator base with EAGLE3 and Qwen3.5 MTP draft models, target-native speculation, and method-keyed draft weight registry - Block speculative verification: paged CuTe verification attention, SM90 GDR verify and commit kernels with persistent pipelining - Unified parameterized batch-op pipeline; executor forward split into target pass and speculative round; host batch ops moved engine-side - Scheduler: unified required admission, simplified rollback, submitted rows carry producer-set effects instead of a kind tag - Op-level pytest suites: verification attention, draft carry, speculative sampling, speculative sequence, target hidden projection, copy --- lmdeploy/metrics/stats.py | 23 +- lmdeploy/serve/core/async_engine.py | 15 +- lmdeploy/turbomind/__init__.py | 5 + lmdeploy/turbomind/builders/__init__.py | 6 + lmdeploy/turbomind/builders/eagle3_weight.py | 13 + .../turbomind/builders/qwen3_5_mtp_weight.py | 12 + lmdeploy/turbomind/builders/text_model.py | 25 +- lmdeploy/turbomind/model_loader.py | 77 +- lmdeploy/turbomind/models/__init__.py | 1 + lmdeploy/turbomind/models/qwen3_5_mtp.py | 66 + lmdeploy/turbomind/models/qwen3_eagle3.py | 76 + lmdeploy/turbomind/spec_decode.py | 68 + lmdeploy/turbomind/text_model.py | 4 +- lmdeploy/turbomind/turbomind.py | 188 +- scripts/test_turbomind_model.py | 155 +- src/turbomind/comm/CMakeLists.txt | 2 + src/turbomind/comm/cuda_ipc/allgather.cu | 5 + .../comm/cuda_ipc/fused_allreduce_ex.cu | 30 +- src/turbomind/comm/nccl/nccl.cu | 59 +- src/turbomind/comm/padded_row_allgather.cc | 52 + src/turbomind/comm/padded_row_allgather.h | 20 + src/turbomind/comm/token_ownership.h | 69 + src/turbomind/engine/CMakeLists.txt | 1 + src/turbomind/engine/README.md | 641 ++++- src/turbomind/engine/batch.h | 5 + src/turbomind/engine/engine.cc | 352 ++- src/turbomind/engine/engine.h | 21 +- src/turbomind/engine/engine_config.h | 3 + src/turbomind/engine/gateway.cc | 6 +- src/turbomind/engine/model.cc | 60 + src/turbomind/engine/model.h | 68 + src/turbomind/engine/model_executor.cc | 417 ++- src/turbomind/engine/model_executor.h | 10 +- src/turbomind/engine/model_request.cc | 8 +- src/turbomind/engine/model_request.h | 4 +- src/turbomind/engine/request.cc | 36 +- src/turbomind/engine/request.h | 87 +- src/turbomind/engine/scheduler.cc | 275 +- src/turbomind/engine/scheduler.h | 7 +- src/turbomind/generation/CMakeLists.txt | 4 + src/turbomind/generation/generation.cc | 424 ++-- src/turbomind/generation/generation.h | 18 +- src/turbomind/generation/generation_impl.h | 124 + src/turbomind/generation/guided_decoding.cc | 2 +- src/turbomind/generation/logits_processor.cc | 93 +- src/turbomind/generation/logits_processor.h | 11 +- src/turbomind/generation/sampling.cc | 148 +- src/turbomind/generation/sampling.h | 9 +- src/turbomind/generation/stop_criteria.cc | 63 +- src/turbomind/generation/stop_criteria.h | 7 + .../generation/target_verification.cc | 245 ++ .../generation/target_verification.h | 82 + src/turbomind/kernels/CMakeLists.txt | 28 + .../kernels/attention/CMakeLists.txt | 2 + src/turbomind/kernels/attention/cp_utils.cu | 22 + src/turbomind/kernels/attention/cp_utils.h | 2 + .../kernels/attention/kv_cache_utils_v2.cu | 16 + .../kernels/attention/kv_cache_utils_v2.h | 4 + .../kernels/attention/rotary_embedding.h | 170 +- .../kernels/attention/test_attention.cu | 3 + .../attention/verification/CMakeLists.txt | 26 + .../attention/verification/attention.h | 125 + .../attention/verification/dispatch.cu | 110 + .../attention/verification/kernel_sm80.cuh | 499 ++++ .../verification/kernel_sm90_wgmma.cuh | 659 +++++ .../verification/kernel_sm90_wgmma_rs.cuh | 776 ++++++ .../verification/kernel_sm90_wgmma_rs_ws.cuh | 790 ++++++ .../attention/verification/paged_kv.cuh | 161 ++ .../attention/verification/policy_sm90.cuh | 159 ++ .../attention/verification/python_bind.cpp | 365 +++ .../kernels/attention/verification/reduce.cu | 136 + .../kernels/attention/verification/stub.cc | 21 + src/turbomind/kernels/ban_bad_words.cu | 7 + src/turbomind/kernels/ban_bad_words.h | 1 + src/turbomind/kernels/copy/copy.cc | 125 +- src/turbomind/kernels/copy/copy.cu | 51 +- src/turbomind/kernels/copy/copy.h | 3 + src/turbomind/kernels/copy/transpose.cu | 52 +- src/turbomind/kernels/draft_carry_kernels.cu | 81 + src/turbomind/kernels/draft_carry_kernels.h | 21 + .../kernels/draft_carry_python_bind.cpp | 64 + src/turbomind/kernels/gemm/CMakeLists.txt | 22 +- .../kernels/linear_attn/CMakeLists.txt | 1 + .../kernels/linear_attn/delta_rule.cu | 16 +- .../kernels/linear_attn/delta_rule.h | 25 +- .../linear_attn/gdn_state_transaction.cu | 391 +++ .../linear_attn/gdn_state_transaction.h | 73 + .../kernels/linear_attn/kernel/CMakeLists.txt | 3 +- .../kernels/linear_attn/kernel/plan.cc | 13 +- .../linear_attn/kernel/sm_90/internal.h | 2 + .../kernels/linear_attn/kernel/sm_90/plan.cc | 2 +- .../linear_attn/kernel/sm_90/recurrent.cu | 39 +- .../kernel/sm_90/verification_fwd.cu | 2261 +++++++++++++++++ .../kernels/linear_attn/python_bind.cpp | 191 +- src/turbomind/kernels/norm/rms_norm.cu | 24 + src/turbomind/kernels/norm/rms_norm.h | 11 + src/turbomind/kernels/sampling_device.cuh | 55 + src/turbomind/kernels/sampling_kernels.cu | 71 +- src/turbomind/kernels/sampling_kernels.h | 10 +- .../kernels/sampling_penalty_kernels.cu | 15 +- .../kernels/sampling_penalty_kernels.h | 1 + .../kernels/sampling_topp_kernels.cu | 51 +- src/turbomind/kernels/sampling_topp_kernels.h | 4 + .../kernels/speculative_sampling_kernels.cu | 308 +++ .../kernels/speculative_sampling_kernels.h | 39 + .../speculative_sampling_python_bind.cpp | 240 ++ .../kernels/speculative_sequence_kernels.cu | 489 ++++ .../kernels/speculative_sequence_kernels.h | 74 + .../speculative_sequence_python_bind.cpp | 295 +++ .../kernels/stop_criteria_kernels.cu | 187 +- src/turbomind/kernels/stop_criteria_kernels.h | 31 +- src/turbomind/models/CMakeLists.txt | 23 + src/turbomind/models/batch_status.cc | 172 ++ src/turbomind/models/batch_status.h | 49 + src/turbomind/models/input_processor.cc | 465 ++-- src/turbomind/models/input_processor.h | 15 +- src/turbomind/models/input_processor_impl.h | 88 + .../models/input_processor_speculative.cc | 104 + src/turbomind/models/internvit/internvit.cc | 62 +- src/turbomind/models/internvit/internvit.h | 6 +- src/turbomind/models/language_model.cc | 669 ++--- src/turbomind/models/language_model.h | 36 + .../models/llama/GatedDeltaNetLayer.cc | 433 +++- .../models/llama/GatedDeltaNetLayer.h | 26 + src/turbomind/models/llama/LlamaFfnLayer.cc | 1 + .../models/llama/context_token_resource.h | 29 +- src/turbomind/models/llama/llama_kernels.cu | 40 +- src/turbomind/models/llama/llama_kernels.h | 7 +- src/turbomind/models/llama/llama_rope.h | 3 +- src/turbomind/models/llama/llama_utils.h | 13 - .../models/llama/unified_attention_layer.cc | 296 ++- .../models/llama/unified_attention_layer.h | 22 + src/turbomind/models/llama/unified_decoder.cc | 112 +- src/turbomind/models/llama/unified_decoder.h | 31 +- src/turbomind/models/model_root.h | 6 + src/turbomind/models/model_weight.cc | 22 +- src/turbomind/models/model_weight.h | 5 +- src/turbomind/models/output_processor.cc | 47 +- src/turbomind/models/output_processor.h | 5 +- src/turbomind/models/qwenvit/qwenvit.cc | 80 +- src/turbomind/models/qwenvit/qwenvit.h | 6 +- .../speculative/collect_hidden_states.cc | 82 + .../speculative/collect_hidden_states.h | 73 + .../models/speculative/eagle3/eagle3_model.cc | 159 ++ .../models/speculative/eagle3/eagle3_model.h | 48 + .../speculative/eagle3/eagle3_weight.cc | 32 + .../models/speculative/eagle3/eagle3_weight.h | 54 + .../eagle3/target_hidden_projection.cc | 117 + .../eagle3/target_hidden_projection.h | 70 + .../target_hidden_projection_kernels.cu | 35 + .../eagle3/target_hidden_projection_kernels.h | 20 + .../target_hidden_projection_python_bind.cpp | 59 + .../models/speculative/fixed_chain_model.cc | 233 ++ .../models/speculative/fixed_chain_model.h | 100 + .../models/speculative/fixed_chain_policy.cc | 27 + .../models/speculative/fixed_chain_policy.h | 26 + .../models/speculative/fixed_chain_setup.cc | 90 + .../models/speculative/fixed_chain_setup.h | 51 + .../models/speculative/hidden_state_tap.h | 54 + .../qwen3_5_mtp/qwen3_5_mtp_model.cc | 179 ++ .../qwen3_5_mtp/qwen3_5_mtp_model.h | 64 + .../qwen3_5_mtp/qwen3_5_mtp_weight.cc | 12 + .../qwen3_5_mtp/qwen3_5_mtp_weight.h | 45 + .../qwen3_5_mtp/target_final_hidden.cc | 66 + .../qwen3_5_mtp/target_final_hidden.h | 45 + src/turbomind/models/speculative/registry.cc | 33 + src/turbomind/models/speculative/registry.h | 40 + .../models/speculative/speculative_model.h | 66 + .../models/speculative/speculative_policy.h | 58 + src/turbomind/models/vision_model.cc | 9 +- src/turbomind/models/vision_model.h | 28 +- src/turbomind/python/CMakeLists.txt | 17 +- .../python/attention_component_bindings.h | 9 + src/turbomind/python/bind.cpp | 92 +- .../python/eagle3_component_bindings.h | 15 + src/turbomind/python/eagle3_dlpack_internal.h | 147 ++ src/turbomind/turbomind.cc | 55 +- src/turbomind/utils/metrics.h | 15 + src/turbomind/utils/nvtx_utils.cc | 21 +- src/turbomind/utils/nvtx_utils.h | 31 +- tests/turbomind/attention/__init__.py | 0 .../attention/test_verification_attention.py | 380 +++ .../attention/verification_attention.py | 86 + tests/turbomind/copy/test_copy.py | 144 ++ tests/turbomind/draft_carry/__init__.py | 1 + tests/turbomind/draft_carry/draft_carry.py | 48 + tests/turbomind/draft_carry/reference.py | 61 + .../turbomind/draft_carry/test_draft_carry.py | 326 +++ tests/turbomind/linear_attn/benchmark.py | 21 +- .../linear_attn/test_gated_delta_rule.py | 440 ++++ .../linear_attn/turbomind_gated_delta_rule.py | 147 +- .../speculative_sampling/__init__.py | 1 + .../speculative_sampling/reference.py | 52 + .../speculative_sampling.py | 108 + .../test_speculative_sampling.py | 875 +++++++ .../speculative_sequence/__init__.py | 1 + .../speculative_sequence/reference.py | 273 ++ .../speculative_sequence.py | 157 ++ .../test_speculative_sequence.py | 1581 ++++++++++++ .../target_hidden_projection/__init__.py | 0 .../target_hidden_projection/reference.py | 17 + .../target_hidden_projection.py | 46 + .../test_target_hidden_projection.py | 330 +++ 203 files changed, 21841 insertions(+), 2129 deletions(-) create mode 100644 lmdeploy/turbomind/builders/eagle3_weight.py create mode 100644 lmdeploy/turbomind/builders/qwen3_5_mtp_weight.py create mode 100644 lmdeploy/turbomind/models/qwen3_5_mtp.py create mode 100644 lmdeploy/turbomind/models/qwen3_eagle3.py create mode 100644 lmdeploy/turbomind/spec_decode.py create mode 100644 src/turbomind/comm/padded_row_allgather.cc create mode 100644 src/turbomind/comm/padded_row_allgather.h create mode 100644 src/turbomind/comm/token_ownership.h create mode 100644 src/turbomind/engine/model.cc create mode 100644 src/turbomind/engine/model.h create mode 100644 src/turbomind/generation/generation_impl.h create mode 100644 src/turbomind/generation/target_verification.cc create mode 100644 src/turbomind/generation/target_verification.h create mode 100644 src/turbomind/kernels/attention/verification/CMakeLists.txt create mode 100644 src/turbomind/kernels/attention/verification/attention.h create mode 100644 src/turbomind/kernels/attention/verification/dispatch.cu create mode 100644 src/turbomind/kernels/attention/verification/kernel_sm80.cuh create mode 100644 src/turbomind/kernels/attention/verification/kernel_sm90_wgmma.cuh create mode 100644 src/turbomind/kernels/attention/verification/kernel_sm90_wgmma_rs.cuh create mode 100644 src/turbomind/kernels/attention/verification/kernel_sm90_wgmma_rs_ws.cuh create mode 100644 src/turbomind/kernels/attention/verification/paged_kv.cuh create mode 100644 src/turbomind/kernels/attention/verification/policy_sm90.cuh create mode 100644 src/turbomind/kernels/attention/verification/python_bind.cpp create mode 100644 src/turbomind/kernels/attention/verification/reduce.cu create mode 100644 src/turbomind/kernels/attention/verification/stub.cc create mode 100644 src/turbomind/kernels/draft_carry_kernels.cu create mode 100644 src/turbomind/kernels/draft_carry_kernels.h create mode 100644 src/turbomind/kernels/draft_carry_python_bind.cpp create mode 100644 src/turbomind/kernels/linear_attn/gdn_state_transaction.cu create mode 100644 src/turbomind/kernels/linear_attn/gdn_state_transaction.h create mode 100644 src/turbomind/kernels/linear_attn/kernel/sm_90/verification_fwd.cu create mode 100644 src/turbomind/kernels/sampling_device.cuh create mode 100644 src/turbomind/kernels/speculative_sampling_kernels.cu create mode 100644 src/turbomind/kernels/speculative_sampling_kernels.h create mode 100644 src/turbomind/kernels/speculative_sampling_python_bind.cpp create mode 100644 src/turbomind/kernels/speculative_sequence_kernels.cu create mode 100644 src/turbomind/kernels/speculative_sequence_kernels.h create mode 100644 src/turbomind/kernels/speculative_sequence_python_bind.cpp create mode 100644 src/turbomind/models/batch_status.cc create mode 100644 src/turbomind/models/batch_status.h create mode 100644 src/turbomind/models/input_processor_impl.h create mode 100644 src/turbomind/models/input_processor_speculative.cc create mode 100644 src/turbomind/models/speculative/collect_hidden_states.cc create mode 100644 src/turbomind/models/speculative/collect_hidden_states.h create mode 100644 src/turbomind/models/speculative/eagle3/eagle3_model.cc create mode 100644 src/turbomind/models/speculative/eagle3/eagle3_model.h create mode 100644 src/turbomind/models/speculative/eagle3/eagle3_weight.cc create mode 100644 src/turbomind/models/speculative/eagle3/eagle3_weight.h create mode 100644 src/turbomind/models/speculative/eagle3/target_hidden_projection.cc create mode 100644 src/turbomind/models/speculative/eagle3/target_hidden_projection.h create mode 100644 src/turbomind/models/speculative/eagle3/target_hidden_projection_kernels.cu create mode 100644 src/turbomind/models/speculative/eagle3/target_hidden_projection_kernels.h create mode 100644 src/turbomind/models/speculative/eagle3/target_hidden_projection_python_bind.cpp create mode 100644 src/turbomind/models/speculative/fixed_chain_model.cc create mode 100644 src/turbomind/models/speculative/fixed_chain_model.h create mode 100644 src/turbomind/models/speculative/fixed_chain_policy.cc create mode 100644 src/turbomind/models/speculative/fixed_chain_policy.h create mode 100644 src/turbomind/models/speculative/fixed_chain_setup.cc create mode 100644 src/turbomind/models/speculative/fixed_chain_setup.h create mode 100644 src/turbomind/models/speculative/hidden_state_tap.h create mode 100644 src/turbomind/models/speculative/qwen3_5_mtp/qwen3_5_mtp_model.cc create mode 100644 src/turbomind/models/speculative/qwen3_5_mtp/qwen3_5_mtp_model.h create mode 100644 src/turbomind/models/speculative/qwen3_5_mtp/qwen3_5_mtp_weight.cc create mode 100644 src/turbomind/models/speculative/qwen3_5_mtp/qwen3_5_mtp_weight.h create mode 100644 src/turbomind/models/speculative/qwen3_5_mtp/target_final_hidden.cc create mode 100644 src/turbomind/models/speculative/qwen3_5_mtp/target_final_hidden.h create mode 100644 src/turbomind/models/speculative/registry.cc create mode 100644 src/turbomind/models/speculative/registry.h create mode 100644 src/turbomind/models/speculative/speculative_model.h create mode 100644 src/turbomind/models/speculative/speculative_policy.h create mode 100644 src/turbomind/python/attention_component_bindings.h create mode 100644 src/turbomind/python/eagle3_component_bindings.h create mode 100644 src/turbomind/python/eagle3_dlpack_internal.h create mode 100644 tests/turbomind/attention/__init__.py create mode 100644 tests/turbomind/attention/test_verification_attention.py create mode 100644 tests/turbomind/attention/verification_attention.py create mode 100644 tests/turbomind/copy/test_copy.py create mode 100644 tests/turbomind/draft_carry/__init__.py create mode 100644 tests/turbomind/draft_carry/draft_carry.py create mode 100644 tests/turbomind/draft_carry/reference.py create mode 100644 tests/turbomind/draft_carry/test_draft_carry.py create mode 100644 tests/turbomind/speculative_sampling/__init__.py create mode 100644 tests/turbomind/speculative_sampling/reference.py create mode 100644 tests/turbomind/speculative_sampling/speculative_sampling.py create mode 100644 tests/turbomind/speculative_sampling/test_speculative_sampling.py create mode 100644 tests/turbomind/speculative_sequence/__init__.py create mode 100644 tests/turbomind/speculative_sequence/reference.py create mode 100644 tests/turbomind/speculative_sequence/speculative_sequence.py create mode 100644 tests/turbomind/speculative_sequence/test_speculative_sequence.py create mode 100644 tests/turbomind/target_hidden_projection/__init__.py create mode 100644 tests/turbomind/target_hidden_projection/reference.py create mode 100644 tests/turbomind/target_hidden_projection/target_hidden_projection.py create mode 100644 tests/turbomind/target_hidden_projection/test_target_hidden_projection.py diff --git a/lmdeploy/metrics/stats.py b/lmdeploy/metrics/stats.py index 68c3798657..5b79e5381e 100644 --- a/lmdeploy/metrics/stats.py +++ b/lmdeploy/metrics/stats.py @@ -285,10 +285,25 @@ def update_from_output(self, outputs: EngineOutput): """Update from engine output.""" spec_info = getattr(outputs.req_metrics, 'spec_info', None) if spec_info: - self.num_drafts += 1 - self.num_draft_tokens += spec_info['num_draft_tokens'] - self.num_accepted_tokens += spec_info['num_accepted_tokens'] - self.num_accepted_tokens_per_pos[:spec_info['num_accepted_tokens']] += 1 + if 'num_drafts' in spec_info: + self.num_drafts += \ + spec_info['num_drafts'] + self.num_draft_tokens += \ + spec_info['num_draft_tokens'] + self.num_accepted_tokens += \ + spec_info['num_accepted_tokens'] + self.num_accepted_tokens_per_pos += \ + np.asarray( + spec_info[ + 'num_accepted_tokens_per_pos']) + else: + self.num_drafts += 1 + self.num_draft_tokens += \ + spec_info['num_draft_tokens'] + self.num_accepted_tokens += \ + spec_info['num_accepted_tokens'] + self.num_accepted_tokens_per_pos[ + :spec_info['num_accepted_tokens']] += 1 def update_per_draft(self, num_draft_tokens: int, num_accepted_tokens: int): """Update with per draft stats.""" diff --git a/lmdeploy/serve/core/async_engine.py b/lmdeploy/serve/core/async_engine.py index 0bff779398..26ac0fd34b 100644 --- a/lmdeploy/serve/core/async_engine.py +++ b/lmdeploy/serve/core/async_engine.py @@ -145,13 +145,12 @@ def __init__(self, self.session_len = (_get_and_verify_max_len(self.hf_cfg, None) if backend_config.session_len is None else backend_config.session_len) backend_config.session_len = self.session_len - if speculative_config is not None and backend == 'turbomind': - logger.warning('speculative decoding is not supported by turbomind ') # build backend engine if backend == 'turbomind': self.engine = self._build_turbomind(model_path=model_path, backend_config=backend_config, trust_remote_code=trust_remote_code, + speculative_config=speculative_config, **kwargs) elif backend == 'pytorch': self.engine = self._build_pytorch(model_path=model_path, @@ -174,8 +173,9 @@ def __init__(self, self.backend = backend self.request_logger = RequestLogger(max_log_len) - self.num_spec_token = 0 if backend == 'turbomind' or speculative_config is None \ - else speculative_config.num_speculative_tokens + self.num_spec_token = (0 + if speculative_config is None + else speculative_config.num_speculative_tokens) self.session_mgr = SessionManager() self.session_mgr.build_request_handle_pool(self.engine, self.backend_config.max_batch_size) @@ -201,6 +201,7 @@ def __exit__(self, exc_type, exc_value, traceback): def _build_turbomind(self, model_path: str, backend_config: TurbomindEngineConfig | None = None, + speculative_config: SpeculativeConfig | None = None, trust_remote_code: bool = False, **kwargs): """Inner build method for turbomind backend.""" @@ -210,7 +211,11 @@ def _build_turbomind(self, 'TurboMind was requested but its native module is unavailable.' ) from turbomind._import_error return turbomind.TurboMind.from_pretrained( - model_path, engine_config=backend_config, trust_remote_code=trust_remote_code, **kwargs + model_path, + engine_config=backend_config, + trust_remote_code=trust_remote_code, + speculative_config=speculative_config, + **kwargs ) def _build_pytorch(self, diff --git a/lmdeploy/turbomind/__init__.py b/lmdeploy/turbomind/__init__.py index 8c9126d3d0..1ae39755bc 100644 --- a/lmdeploy/turbomind/__init__.py +++ b/lmdeploy/turbomind/__init__.py @@ -1,5 +1,7 @@ # Copyright (c) OpenMMLab. All rights reserved. +import sys + import torch # noqa: F401 _import_error = None @@ -10,6 +12,9 @@ _tm = None _import_error = error else: + # Expose the extension under its bare name so submodules that predate the + # lazy-import design can `import _turbomind` without a second load. + sys.modules.setdefault('_turbomind', _tm) from .turbomind import TurboMind as TurboMind diff --git a/lmdeploy/turbomind/builders/__init__.py b/lmdeploy/turbomind/builders/__init__.py index 63a00f34cd..af67a32e7e 100644 --- a/lmdeploy/turbomind/builders/__init__.py +++ b/lmdeploy/turbomind/builders/__init__.py @@ -7,11 +7,13 @@ from .attention import AttentionBuilder from .decoder_layer import DecoderLayerBuilder, DecoderLayerConfig from .deltanet import DeltaNetBuilder +from .eagle3_weight import Eagle3WeightBuilder, Eagle3WeightConfig from .ffn import FfnBuilder from .mla import MLABuilder from .module_list import ModuleListBuilder, ModuleListConfig from .moe import MoeBuilder from .norm import LayerNormBuilder, NormBuilder, make_layer_norm_config, make_norm_config +from .qwen3_5_mtp_weight import Qwen35MtpWeightBuilder, Qwen35MtpWeightConfig from .text_model import TextModelBuilder from .vision_model import VisionModelBuilder @@ -32,6 +34,8 @@ 'DeltaNetBuilder', 'MLABuilder', 'DecoderLayerBuilder', + 'Eagle3WeightBuilder', + 'Qwen35MtpWeightBuilder', 'ModuleListBuilder', 'NormBuilder', 'LayerNormBuilder', @@ -40,6 +44,8 @@ 'make_layer_norm_config', # C++ config re-exports 'DecoderLayerConfig', + 'Eagle3WeightConfig', + 'Qwen35MtpWeightConfig', 'ModuleListConfig', # Helper functions ] diff --git a/lmdeploy/turbomind/builders/eagle3_weight.py b/lmdeploy/turbomind/builders/eagle3_weight.py new file mode 100644 index 0000000000..35b721f107 --- /dev/null +++ b/lmdeploy/turbomind/builders/eagle3_weight.py @@ -0,0 +1,13 @@ +# Copyright (c) OpenMMLab. All rights reserved. +import _turbomind as _tm + +from ._base import Builder + +Eagle3WeightConfig = _tm.Eagle3WeightConfig + + +class Eagle3WeightBuilder(Builder): + """Builder for the EAGLE3 method weight tree at ModelWeight.spec.""" + + def add_target_hidden_proj(self, linear): + self._add_linear('target_hidden_proj', linear, split_side=None) diff --git a/lmdeploy/turbomind/builders/qwen3_5_mtp_weight.py b/lmdeploy/turbomind/builders/qwen3_5_mtp_weight.py new file mode 100644 index 0000000000..be8d3508af --- /dev/null +++ b/lmdeploy/turbomind/builders/qwen3_5_mtp_weight.py @@ -0,0 +1,12 @@ +# Copyright (c) OpenMMLab. All rights reserved. +import _turbomind as _tm + +from ._base import Builder, SplitSide + +Qwen35MtpWeightConfig = _tm.Qwen35MtpWeightConfig + + +class Qwen35MtpWeightBuilder(Builder): + + def add_fc(self, linear): + self._add_linear('fc', linear, split_side=SplitSide.OUTPUT) diff --git a/lmdeploy/turbomind/builders/text_model.py b/lmdeploy/turbomind/builders/text_model.py index 8f4b82d893..0bdeeba5bb 100644 --- a/lmdeploy/turbomind/builders/text_model.py +++ b/lmdeploy/turbomind/builders/text_model.py @@ -9,31 +9,38 @@ class TextModelBuilder(Builder): Constructs a ModelWeight via ``_tm.create_module(ModelWeightConfig)`` on each context (inherited Builder machinery), then attaches it to - externally-owned ``ModelRoot`` sentinel handles as their - ``text_model`` child during ``build()``. + externally-owned ``ModelRoot`` sentinel handles as the configured + root child during ``build()``. Owns ``tok_embeddings`` (Tensor param) and ``output`` (LinearWeight child) commits on the ModelWeight via ``add_token_embeds`` / ``add_lm_head``. """ - def __init__(self, config, ctx, *, root_handles, - tp: ParallelGroup, vocab_size): + def __init__(self, + config, + ctx, + *, + root_handles, + tp: ParallelGroup, + vocab_size, + root_child: str = 'text_model'): super().__init__(config, ctx) self.tp = tp self.config.tp_size = tp.size self._root_handles = root_handles self._vocab_size = vocab_size + self._root_child = root_child def build(self) -> BuiltModule: - """Create ModelWeight via _tm.create_module (via super), then attach - each per-GPU ModelWeight handle to its sentinel root via - add_child_raw.""" + """Build and attach each ModelWeight to the configured root child.""" built = super().build() - for i, (root, text_model) in enumerate( + for i, (root, model_weight) in enumerate( zip(self._root_handles, built.handles)): with self._ctx.devices[i]: - root.add_child_raw('text_model', text_model) + root.add_child_raw( + self._root_child, + model_weight) return built def add_token_embeds(self, tensor): diff --git a/lmdeploy/turbomind/model_loader.py b/lmdeploy/turbomind/model_loader.py index 3917c4cbe7..1b25fa3840 100644 --- a/lmdeploy/turbomind/model_loader.py +++ b/lmdeploy/turbomind/model_loader.py @@ -15,12 +15,22 @@ class ModelLoader: to the C++ runtime. """ - def __init__(self, model, model_comm, gpu_count, model_path, - data_type, engine_config): + def __init__(self, + model, + model_comm, + gpu_count, + model_path, + data_type, + engine_config, + *, + draft_model=None, + draft_model_path=None): self.model = model + self.draft_model = draft_model self.model_comm = model_comm self.gpu_count = gpu_count self.model_path = model_path + self.draft_model_path = draft_model_path self.data_type = data_type self.engine_config = engine_config self._bind_runtime() @@ -59,33 +69,52 @@ def _bind_runtime(self): [(e * mlp_tp.size + m) % dense_size for e, m in zip(ep.ranks, mlp_tp.ranks)]) - self.model.bind_runtime( - ctx=ctx, - root_handles=[mc.root(g) for g in range(self.gpu_count)], - attn_tp=attn_tp, - mlp_tp=mlp_tp, - ep=ep, - model_tp=model_tp, - dense_tp=dense_tp, - ) + models = ( + (self.model,) + if self.draft_model is None + else (self.model, self.draft_model)) + for model in models: + model.bind_runtime( + ctx=ctx, + root_handles=[ + mc.root(g) + for g in range(self.gpu_count)], + attn_tp=attn_tp, + mlp_tp=mlp_tp, + ep=ep, + model_tp=model_tp, + dense_tp=dense_tp) - def export(self): - ckpt = create_checkpoint( - self.model_path, - mappings=getattr(self.model, '_loader_mappings', [])) + @staticmethod + def _export_one(model, model_path): + checkpoint = create_checkpoint( + model_path, + mappings=getattr( + model, '_loader_mappings', [])) try: - self.model.model(Prefix(ckpt)) + root = Prefix(checkpoint) + if getattr(model, 'prefix', ''): + root = root + model.prefix + model.model(root) finally: - ckpt.close() + checkpoint.close() + + def export(self): + self._export_one( + self.model, self.model_path) + if self.draft_model is not None: + self._export_one( + self.draft_model, + self.draft_model_path) torch.cuda.empty_cache() def export_iter(self): - ckpt = create_checkpoint( - self.model_path, - mappings=getattr(self.model, '_loader_mappings', [])) - try: - self.model.model(Prefix(ckpt)) + self._export_one( + self.model, self.model_path) + yield -1 + if self.draft_model is not None: + self._export_one( + self.draft_model, + self.draft_model_path) yield -1 - finally: - ckpt.close() torch.cuda.empty_cache() diff --git a/lmdeploy/turbomind/models/__init__.py b/lmdeploy/turbomind/models/__init__.py index 7d873ef619..01c47bc12a 100644 --- a/lmdeploy/turbomind/models/__init__.py +++ b/lmdeploy/turbomind/models/__init__.py @@ -10,3 +10,4 @@ from .qwen2_vl import Qwen2VLModel # noqa: F401 from .qwen3 import Qwen3TextModel # noqa: F401 from .qwen3_5 import Qwen3_5Model, Qwen3_5TextModel, Qwen3_5VisionModel # noqa: F401 +from .qwen3_eagle3 import Qwen3Eagle3TextModel # noqa: F401 diff --git a/lmdeploy/turbomind/models/qwen3_5_mtp.py b/lmdeploy/turbomind/models/qwen3_5_mtp.py new file mode 100644 index 0000000000..b7bb8031e1 --- /dev/null +++ b/lmdeploy/turbomind/models/qwen3_5_mtp.py @@ -0,0 +1,66 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""Checkpoint-native Qwen3.5 MTP draft weight tree.""" + +from ..builders import ( + DecoderLayerBuilder, + DecoderLayerConfig, + ModuleListBuilder, + ModuleListConfig, + Qwen35MtpWeightBuilder, + Qwen35MtpWeightConfig, + TextModelBuilder, +) +from .qwen3_5 import Qwen3_5TextModel +from .utils import make_model_weight_config + + +class Qwen3_5MtpTextModel(Qwen3_5TextModel): + + def __init__(self, cfg, *, resolver): + super().__init__(cfg, resolver=resolver) + self.tap_layer_ids = [cfg.num_hidden_layers] + + def model(self, pfx): + root_cfg = make_model_weight_config(self.cfg) + root_cfg.decoder_only = True + + root = TextModelBuilder( + root_cfg, + self._ctx, + root_handles=self._root_handles, + tp=self._model_tp, + vocab_size=self.cfg.vocab_size, + root_child='draft_model') + + root.norm = self.norm(pfx + 'mtp.norm', zero_centered=True) + root.layers = self.mtp_layers(pfx + 'mtp.layers') + root.spec = self.mtp_spec(pfx + 'mtp') + root.build() + + def mtp_layers(self, pfx): + layers = ModuleListBuilder(ModuleListConfig(), self._ctx) + for i, layer_pfx in pfx.slices(0, 1): + layer = DecoderLayerBuilder(DecoderLayerConfig(), self._ctx) + layer.attention = self.attn(layer_pfx + 'self_attn') + if self._n_experts > 0: + layer.moe_ffn = self.moe(layer_pfx + 'mlp') + else: + layer.feed_forward = self.ffn(layer_pfx + 'mlp', + self.cfg.intermediate_size, + tp=self._mlp_tp) + layer.attention_norm = self.norm( + layer_pfx + 'input_layernorm', zero_centered=True) + layer.ffn_norm = self.norm( + layer_pfx + 'post_attention_layernorm', zero_centered=True) + layers[i] = layer.build() + return layers.build() + + def mtp_spec(self, pfx): + spec = Qwen35MtpWeightBuilder(Qwen35MtpWeightConfig(), self._ctx) + spec.tp = self._model_tp + spec.add_fc(self._linear(pfx + 'fc')) + spec.pre_fc_norm_embedding = self.norm( + pfx + 'pre_fc_norm_embedding', zero_centered=True) + spec.pre_fc_norm_hidden = self.norm( + pfx + 'pre_fc_norm_hidden', zero_centered=True) + return spec.build() diff --git a/lmdeploy/turbomind/models/qwen3_eagle3.py b/lmdeploy/turbomind/models/qwen3_eagle3.py new file mode 100644 index 0000000000..09e4840476 --- /dev/null +++ b/lmdeploy/turbomind/models/qwen3_eagle3.py @@ -0,0 +1,76 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""EAGLE3 draft weight model for Qwen3 targets.""" + +from ..builders import ( + DecoderLayerBuilder, + DecoderLayerConfig, + Eagle3WeightBuilder, + Eagle3WeightConfig, + ModuleListBuilder, + ModuleListConfig, + TextModelBuilder, +) +from .base import INPUT_MODELS +from .qwen3 import Qwen3TextModel +from .utils import make_model_weight_config + + +@INPUT_MODELS.register_module(name='qwen3-eagle3') +class Qwen3Eagle3TextModel(Qwen3TextModel): + + def __init__(self, cfg, *, resolver, prefix=''): + super().__init__(cfg, resolver=resolver) + self.prefix = prefix + self.tap_layer_ids = list(cfg.target_layer_ids) + + def model(self, pfx): + root_cfg = make_model_weight_config(self.cfg) + builder = TextModelBuilder( + root_cfg, + self._ctx, + root_handles=self._root_handles, + tp=self._model_tp, + vocab_size=self.cfg.vocab_size, + root_child='draft_model') + builder.add_token_embeds( + pfx.get('embed_tokens.weight')) + builder.norm = self.norm(pfx + 'norm') + builder.add_lm_head( + self._linear(pfx + 'lm_head')) + builder.layers = self.layers(pfx + 'layers') + builder.spec = self.spec(pfx) + builder.build() + + def spec(self, pfx): + spec = Eagle3WeightBuilder( + Eagle3WeightConfig(), self._ctx) + spec.add_target_hidden_proj( + self._linear(pfx + 'fc')) + spec.hidden_norms = self.hidden_norms( + pfx + 'layers') + return spec.build() + + def hidden_norms(self, pfx): + norms = ModuleListBuilder( + ModuleListConfig(), self._ctx) + for i, p in pfx.slices( + 0, self.cfg.num_hidden_layers): + norms[i] = self.norm(p + 'hidden_norm') + return norms.build() + + def layers(self, pfx): + layers = ModuleListBuilder( + ModuleListConfig(), self._ctx) + for i, p in pfx.slices( + 0, self.cfg.num_hidden_layers): + layer = DecoderLayerBuilder( + DecoderLayerConfig(), self._ctx) + layer.attention_norm = self.norm( + p + 'input_layernorm') + layer.attention = self.attn( + p + 'self_attn') + layer.ffn_norm = self.norm( + p + 'post_attention_layernorm') + layer.feed_forward = self.ffn(p + 'mlp', tp=self._mlp_tp) + layers[i] = layer.build() + return layers.build() diff --git a/lmdeploy/turbomind/spec_decode.py b/lmdeploy/turbomind/spec_decode.py new file mode 100644 index 0000000000..243eae610e --- /dev/null +++ b/lmdeploy/turbomind/spec_decode.py @@ -0,0 +1,68 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""Speculative method registry for the TurboMind backend.""" +from __future__ import annotations + +import os.path as osp +from dataclasses import dataclass + +from lmdeploy.archs import get_model_arch +from lmdeploy.utils import get_model + +from .models.base import INPUT_MODELS +from .models.qwen3_5_mtp import Qwen3_5MtpTextModel +from .models.utils import source_model_config +from .weight_format import TrivialFormat, WeightFormatResolver + + +@dataclass(frozen=True) +class DraftWeightSpec: + """How one speculative method's draft weights are found and mapped.""" + + input_model: str + weight_source: str + prefix: str = '' + quantized: bool = False + + +DRAFT_WEIGHT_SPECS: dict[str, DraftWeightSpec] = { + 'eagle3': DraftWeightSpec(input_model='qwen3-eagle3', weight_source='sidecar'), +} + + +def build_draft_model(speculative_config, + target_model, + target_model_path, + engine_data_type, + download_dir=None): + """Build the draft weight mapper and resolve the checkpoint it reads.""" + method = speculative_config.method + + if method == 'mtp': + target_text = target_model.text_model + draft = Qwen3_5MtpTextModel( + target_text.cfg, resolver=target_text._resolver) + return draft, target_model_path + + spec = DRAFT_WEIGHT_SPECS.get(method) + if spec is None: + raise ValueError(f'TurboMind does not support speculative method {method!r}; ' + f'supported: {sorted(DRAFT_WEIGHT_SPECS)}') + + if spec.weight_source == 'sidecar': + if not speculative_config.model: + raise ValueError(f'speculative method {method!r} requires a draft model path') + path = speculative_config.model + if not osp.exists(path): + path = get_model(path, download_dir) + else: + path = target_model_path + + _, hf_config = get_model_arch(path) + draft_config = source_model_config(hf_config) + + if spec.quantized: + raise NotImplementedError(f'quantized draft weights are not supported yet ({method!r})') + resolver = WeightFormatResolver(formats=[TrivialFormat(weight_dtype=engine_data_type)]) + + model = INPUT_MODELS.get(spec.input_model)(draft_config, resolver=resolver, prefix=spec.prefix) + return model, path diff --git a/lmdeploy/turbomind/text_model.py b/lmdeploy/turbomind/text_model.py index 641696d1aa..47a750a44d 100644 --- a/lmdeploy/turbomind/text_model.py +++ b/lmdeploy/turbomind/text_model.py @@ -32,9 +32,11 @@ class TextModel(ABC): _loader_mappings: list = [] - def __init__(self, cfg: PretrainedConfig, *, resolver): + def __init__(self, cfg: PretrainedConfig, *, resolver, prefix: str = ''): self.cfg: PretrainedConfig = cfg self._resolver = resolver + self.prefix = prefix + self.tap_layer_ids: list[int] = [] @property def _vocab_size(self) -> int: diff --git a/lmdeploy/turbomind/turbomind.py b/lmdeploy/turbomind/turbomind.py index 70e4fd79f0..acf3b539cb 100644 --- a/lmdeploy/turbomind/turbomind.py +++ b/lmdeploy/turbomind/turbomind.py @@ -15,7 +15,14 @@ import torch from lmdeploy._guided_decoding import compile_response_format -from lmdeploy.messages import EngineOutput, GenerationConfig, ResponseType, ScheduleMetrics, TurbomindEngineConfig +from lmdeploy.messages import ( + EngineOutput, + GenerationConfig, + ResponseType, + ScheduleMetrics, + SpeculativeConfig, + TurbomindEngineConfig, +) from lmdeploy.serve.openai.protocol import UpdateParamsRequest from lmdeploy.tokenizer import Tokenizer from lmdeploy.utils import get_logger, get_max_batch_size, get_model @@ -125,9 +132,11 @@ def __init__(self, chat_template_name: str = None, engine_config: TurbomindEngineConfig = None, trust_remote_code: bool = False, + speculative_config: SpeculativeConfig | None = None, **kwargs): self.model_name = model_name self.chat_template_name = chat_template_name + self.speculative_config = speculative_config _engine_config = copy.deepcopy(engine_config) if _engine_config is None: @@ -158,7 +167,8 @@ def __init__(self, if not osp.exists(model_path): model_path = get_model(model_path, _engine_config.download_dir, _engine_config.revision) self.model_comm, model_loader = self._from_hf(model_path=model_path, engine_config=_engine_config, - trust_remote_code=trust_remote_code) + trust_remote_code=trust_remote_code, + speculative_config=speculative_config) self.source_model = model_loader.model self.is_dummy = self.model_comm.is_dummy_node() self.tokenizer = Tokenizer(model_path, trust_remote_code=trust_remote_code) @@ -215,8 +225,11 @@ def _create_weight_func(device_id): for future in futures: future.result() - def _from_hf(self, model_path: str, engine_config: TurbomindEngineConfig, - trust_remote_code: bool = False): + def _from_hf(self, + model_path: str, + engine_config: TurbomindEngineConfig, + trust_remote_code: bool = False, + speculative_config: SpeculativeConfig | None = None): """Load model which is in hf format.""" assert is_supported(model_path, trust_remote_code=trust_remote_code), ( f'turbomind does not support {model_path}. ' @@ -224,10 +237,26 @@ def _from_hf(self, model_path: str, engine_config: TurbomindEngineConfig, from .converter import get_tm_config from .model_loader import ModelLoader + from .spec_decode import build_draft_model model, model_path, data_type = get_tm_config(model_path, engine_config, trust_remote_code=trust_remote_code) + draft_model = None + draft_model_path = None + spec_method = '' + spec_num_draft_tokens = 0 + spec_tap_layer_ids = [] + + if (speculative_config is not None + and speculative_config.num_speculative_tokens > 0): + draft_model, draft_model_path = build_draft_model( + speculative_config, model, model_path, data_type, + engine_config.download_dir) + spec_method = speculative_config.method + spec_num_draft_tokens = speculative_config.num_speculative_tokens + spec_tap_layer_ids = list(draft_model.tap_layer_ids) + self._vocab_size = model._vocab_size self.engine_config = engine_config @@ -267,6 +296,9 @@ def _from_hf(self, model_path: str, engine_config: TurbomindEngineConfig, ec.node_rank = engine_config.node_rank ec.communicator = engine_config.communicator ec.moe_a2a_backend = engine_config.moe_a2a_backend + ec.spec_method = spec_method + ec.spec_num_draft_tokens = spec_num_draft_tokens + ec.spec_tap_layer_ids = spec_tap_layer_ids logger.info(f'turbomind engine config:\n\n' f'dtype={engine_config.dtype}, state_dtype={state_dtype}, ' @@ -289,6 +321,8 @@ def _from_hf(self, model_path: str, engine_config: TurbomindEngineConfig, model_path=model_path, data_type=data_type, engine_config=engine_config, + draft_model=draft_model, + draft_model_path=draft_model_path, ) return model_comm, model_loader @@ -361,6 +395,7 @@ def from_pretrained(cls, chat_template_name: str = None, engine_config: TurbomindEngineConfig = None, trust_remote_code: bool = False, + speculative_config: SpeculativeConfig | None = None, **kwargs): """LMDeploy's turbomind inference engine. @@ -385,6 +420,7 @@ def from_pretrained(cls, chat_template_name=chat_template_name, engine_config=engine_config, trust_remote_code=trust_remote_code, + speculative_config=speculative_config, **kwargs) def close(self): @@ -523,12 +559,29 @@ def _func(out: EngineOutput, step: int, **kwargs): def _get_metrics(metrics): import time - from lmdeploy.messages import EngineEvent, EventType, RequestMetrics + from lmdeploy.messages import ( + EngineEvent, + EventType, + RequestMetrics, + ) is_first = True - - def _func(out: EngineOutput, step: int, **kwargs): + previous_num_drafts = 0 + previous_num_draft_tokens = 0 + previous_num_accepted_tokens = 0 + previous_num_accepted_tokens_per_pos = None + + def _func( + out: EngineOutput, + step: int, + state=None, + **kwargs): nonlocal is_first + nonlocal previous_num_drafts + nonlocal previous_num_draft_tokens + nonlocal previous_num_accepted_tokens + nonlocal previous_num_accepted_tokens_per_pos + cached_tokens = metrics.cached_tokens if not is_first: out.req_metrics = RequestMetrics(token_timestamp=time.time(), cached_tokens=cached_tokens) @@ -541,6 +594,56 @@ def _func(out: EngineOutput, step: int, **kwargs): cached_tokens=cached_tokens) is_first = False + if state is None: + return + + current_per_pos = list( + state.num_accepted_tokens_per_pos) + if not current_per_pos: + return + + if previous_num_accepted_tokens_per_pos is None: + previous_num_accepted_tokens_per_pos = [ + 0 for _ in current_per_pos + ] + + delta_num_drafts = ( + state.num_drafts + - previous_num_drafts) + delta_num_draft_tokens = ( + state.num_draft_tokens + - previous_num_draft_tokens) + delta_num_accepted_tokens = ( + state.num_accepted_tokens + - previous_num_accepted_tokens) + delta_num_accepted_tokens_per_pos = [ + current - previous + for current, previous in zip( + current_per_pos, + previous_num_accepted_tokens_per_pos) + ] + + previous_num_drafts = state.num_drafts + previous_num_draft_tokens = ( + state.num_draft_tokens) + previous_num_accepted_tokens = ( + state.num_accepted_tokens) + previous_num_accepted_tokens_per_pos = ( + current_per_pos) + + if delta_num_drafts == 0: + return + + out.req_metrics.spec_info = { + 'num_drafts': delta_num_drafts, + 'num_draft_tokens': + delta_num_draft_tokens, + 'num_accepted_tokens': + delta_num_accepted_tokens, + 'num_accepted_tokens_per_pos': + delta_num_accepted_tokens_per_pos, + } + return _func @@ -616,20 +719,32 @@ def _get_extra_output_processors(self, outputs: dict[str, torch.Tensor], gen_con def _get_offset(type): return input_len - 1 if type == 'generation' else 0 - fs = [] + optional_output_fs = [] + if gen_config.output_logits: offset = _get_offset(gen_config.output_logits) - fs.append(_get_logits(outputs, offset)) + optional_output_fs.append( + _get_logits(outputs, offset)) if gen_config.return_ppl: - fs.append(_get_ce_loss(outputs)) + optional_output_fs.append( + _get_ce_loss(outputs)) if gen_config.output_last_hidden_state: - offset = _get_offset(gen_config.output_last_hidden_state) - fs.append(_get_last_hidden_state(outputs, offset)) + offset = _get_offset( + gen_config.output_last_hidden_state) + optional_output_fs.append( + _get_last_hidden_state(outputs, offset)) if gen_config.logprobs: - fs.append(_get_logprobs(outputs, gen_config.logprobs)) - if self.tm_model.engine_config.enable_metrics: - fs.append(_get_metrics(metrics)) - return fs + optional_output_fs.append( + _get_logprobs( + outputs, + gen_config.logprobs)) + + metrics_f = ( + _get_metrics(metrics) + if self.tm_model.engine_config.enable_metrics + else None) + + return optional_output_fs, metrics_f def prepare_embeddings(self, input_embeddings=None, input_embedding_ranges=None): """Convert embeddings.""" @@ -708,22 +823,30 @@ async def async_stream_infer(self, kwargs (dict): kwargs for backward compatibility """ logger.info(f'[async_stream_infer] session {session_id} start') - gen_cfg = self._get_generation_config(gen_config) + local_gen_config = gen_config + if self.tm_model.speculative_config is not None: + local_gen_config = copy.copy(gen_config) + local_gen_config.output_logits = None + local_gen_config.output_last_hidden_state = None + local_gen_config.logprobs = None + local_gen_config.return_ppl = False + + gen_cfg = self._get_generation_config(local_gen_config) inputs, input_len = self.prepare_inputs(input_ids=input_ids, input_embeddings=input_embeddings, input_embedding_ranges=input_embedding_ranges, - gen_config=gen_config) + gen_config=local_gen_config) - if gen_config.response_format is not None: + if local_gen_config.response_format is not None: try: compiler = self.tm_model.grammar_compiler - grammar = compile_response_format(compiler, gen_config.response_format) + grammar = compile_response_format(compiler, local_gen_config.response_format) self.model_inst.set_grammar(grammar) except (ValueError, KeyError) as e: logger.warning(f'Failed to initialize guided decoding, ' f'disable guided decoding: {e}') - gen_config.response_format = None + local_gen_config.response_format = None session = _tm.SessionParam(id=session_id, step=0) @@ -738,7 +861,12 @@ async def async_stream_infer(self, outputs = _tm_dict_to_torch_dict(outputs) - extra_fs = self._get_extra_output_processors(outputs, gen_config, input_len, metrics) + optional_output_fs, metrics_f = ( + self._get_extra_output_processors( + outputs, + local_gen_config, + input_len, + metrics)) output_ids_buf = outputs['output_ids'] @@ -760,7 +888,15 @@ async def async_stream_infer(self, ret_status = ResponseType.FINISH if status == 7 else ResponseType.CANCEL elif status: logger.error(f'internal error. status_code {status}') - yield self._get_error_output(status) + output = self._get_error_output(status) + if (metrics_f is not None + and self.tm_model.speculative_config + is not None): + metrics_f( + output, + seq_len, + state=state) + yield output break if seq_len == prev_len and not finish: @@ -769,8 +905,10 @@ async def async_stream_infer(self, output_ids = output_ids_buf[prev_len:seq_len].tolist() output = EngineOutput(ret_status, output_ids) - for f in extra_fs: - f(output, seq_len) + for f in optional_output_fs: + f(output, seq_len, state=state) + if metrics_f is not None: + metrics_f(output, seq_len, state=state) prev_len = seq_len diff --git a/scripts/test_turbomind_model.py b/scripts/test_turbomind_model.py index 3c17b129b6..04aae34308 100644 --- a/scripts/test_turbomind_model.py +++ b/scripts/test_turbomind_model.py @@ -22,6 +22,9 @@ cache_prompt: 'auto' cache_generation: 'auto' cache_prompt_boundary_skip: 1 + speculative_method: None + speculative_model: '' + num_speculative_tokens: 1 prompt_count: 1 prompt_source: default CUDA_LAUNCH_BLOCKING: 1 (only if --debug was passed) @@ -66,6 +69,9 @@ [--cache-checkpoint-interval N] \\ [--cache-prompt {all,auto}] \\ [--cache-generation {all,auto,none}] \\ + [--speculative-method METHOD] \\ + [--speculative-model MODEL] \\ + [--num-speculative-tokens N] \\ [--debug] Optional prompts: repeat --prompt for multiple strings, or --prompt-file for a JSON @@ -174,10 +180,33 @@ class ResolvedPrompts(NamedTuple): source: str # 'default' | 'cli' | 'file' +class SpeculativeMetrics(NamedTuple): + num_drafts: int + num_draft_tokens: int + num_accepted_tokens: int + num_accepted_tokens_per_pos: list[int] + + @property + def draft_acceptance_rate(self) -> float: + return self.num_accepted_tokens / self.num_draft_tokens + + @property + def mean_acceptance_length(self) -> float: + return 1 + self.num_accepted_tokens / self.num_drafts + + @property + def per_position_acceptance_rate(self) -> list[float]: + return [ + accepted / self.num_drafts + for accepted in self.num_accepted_tokens_per_pos + ] + + class SmokeResult(NamedTuple): create_s: float infer_s: float responses: list[PromptResult] + speculative_metrics: SpeculativeMetrics | None def _set_hf_cache(path: str) -> None: @@ -311,6 +340,9 @@ def build_arg_parser() -> argparse.ArgumentParser: --cache-checkpoint-interval --cache-prompt --cache-generation + --speculative-method + --speculative-model + --num-speculative-tokens Exit 0: load + inference complete. Exit 1: exception (traceback on stderr). Exit 2: usage error. """, ) @@ -396,6 +428,22 @@ def build_arg_parser() -> argparse.ArgumentParser: help=('TurbomindEngineConfig.cache_prompt_boundary_skip ' f'(default: {DEFAULT_CACHE_PROMPT_BOUNDARY_SKIP})'), ) + parser.add_argument( + '--speculative-method', + default=None, + help='Speculative decoding method (default: disabled)', + ) + parser.add_argument( + '--speculative-model', + default='', + help='Speculative draft model id or local path (default: empty)', + ) + parser.add_argument( + '--num-speculative-tokens', + type=int, + default=1, + help='Number of speculative draft tokens (default: 1)', + ) parser.add_argument( '--prompt', action='append', @@ -451,6 +499,9 @@ def run_smoke_infer( cache_prompt: str = DEFAULT_CACHE_PROMPT, cache_generation: str = DEFAULT_CACHE_GENERATION, cache_prompt_boundary_skip: int = DEFAULT_CACHE_PROMPT_BOUNDARY_SKIP, + speculative_method: str | None = None, + speculative_model: str = '', + num_speculative_tokens: int = 1, debug: bool = False, ) -> SmokeResult: _validate_engine_params( @@ -471,6 +522,7 @@ def run_smoke_infer( os.environ['CUDA_LAUNCH_BLOCKING'] = '1' from lmdeploy import GenerationConfig, TurbomindEngineConfig, pipeline + from lmdeploy.messages import SpeculativeConfig engine_config = TurbomindEngineConfig( async_=async_, @@ -482,7 +534,7 @@ def run_smoke_infer( cp=cp, dp=dp, ep=ep, - enable_metrics=False, + enable_metrics=speculative_method is not None, communicator=communicator, enable_prefix_caching=enable_prefix_caching, cache_checkpoint_interval=cache_checkpoint_interval, @@ -491,14 +543,70 @@ def run_smoke_infer( cache_prompt_boundary_skip=cache_prompt_boundary_skip, ) gen_config = GenerationConfig(max_new_tokens=max_new_tokens, do_sample=False) + speculative_config = ( + SpeculativeConfig( + method=speculative_method, + model=speculative_model, + num_speculative_tokens=num_speculative_tokens, + ) if speculative_method is not None else None) + + if speculative_config is not None: + from lmdeploy.metrics import loggers as metrics_loggers + + # The smoke report consumes the in-process logging counters directly; + # it does not create the optional Prometheus exporter. + class _NoopPrometheusStatLogger: + + def __init__(self, model_name, max_model_len, dp_rank): + pass + + def record_schedule(self, stats): + pass + + def record_iteration(self, stats): + pass + + def record_specdecode(self, stats): + pass + + def record_finish(self, stats): + pass + + metrics_loggers.PrometheusStatLogger = ( + _NoopPrometheusStatLogger) + speculative_metrics = None t0 = time.perf_counter() with pipeline(model_id, backend_config=engine_config, log_level='WARNING', - trust_remote_code=True) as pipe: + trust_remote_code=True, speculative_config=speculative_config) as pipe: create_s = time.perf_counter() - t0 - t1 = time.perf_counter() - out = pipe(resolved.prompts, gen_config=gen_config, do_preprocess=True) - infer_s = time.perf_counter() - t1 + metrics_processor = None + if speculative_config is not None: + from lmdeploy.metrics.metrics_processor import metrics_processor + pipe._run( + fn=lambda: metrics_processor.start_metrics_handler( + enable_metrics=True)).result() + try: + t1 = time.perf_counter() + out = pipe(resolved.prompts, gen_config=gen_config, do_preprocess=True) + infer_s = time.perf_counter() - t1 + if metrics_processor is not None: + pipe._run( + coro=metrics_processor.metrics_queue.join()).result() + logger = pipe.async_engine.stat_loggers[0] + speculative_metrics = SpeculativeMetrics( + num_drafts=logger.num_drafts, + num_draft_tokens=logger.num_draft_tokens, + num_accepted_tokens=logger.num_accepted_tokens, + num_accepted_tokens_per_pos=[ + int(x) + for x in logger.num_accepted_tokens_per_pos + ], + ) + finally: + if metrics_processor is not None: + pipe._run( + coro=metrics_processor.stop_metrics_handler()).result() if not isinstance(out, list): out = [out] @@ -518,7 +626,7 @@ def run_smoke_infer( input_token_len=getattr(res, 'input_token_len', -1), generate_token_len=getattr(res, 'generate_token_len', -1), )) - return SmokeResult(create_s, infer_s, responses) + return SmokeResult(create_s, infer_s, responses, speculative_metrics) def print_report( @@ -542,6 +650,9 @@ def print_report( cache_prompt: str = DEFAULT_CACHE_PROMPT, cache_generation: str = DEFAULT_CACHE_GENERATION, cache_prompt_boundary_skip: int = DEFAULT_CACHE_PROMPT_BOUNDARY_SKIP, + speculative_method: str | None = None, + speculative_model: str = '', + num_speculative_tokens: int = 1, debug: bool = False, ) -> None: print('--- setup ---') @@ -562,6 +673,9 @@ def print_report( print(f'cache_generation: {cache_generation!r}') print(f'cache_prompt_boundary_skip: {cache_prompt_boundary_skip}') print(f'max_prefill_token_num: {max_prefill_token_num}') + print(f'speculative_method: {speculative_method!r}') + print(f'speculative_model: {speculative_model!r}') + print(f'num_speculative_tokens: {num_speculative_tokens}') print(f'prompt_count: {len(resolved.prompts)}') print(f'prompt_source: {resolved.source}') if debug: @@ -570,6 +684,23 @@ def print_report( print('--- timing ---') print(f'pipeline load: {result.create_s:.2f} s') print(f'inference: {result.infer_s:.2f} s') + if result.speculative_metrics is not None: + metrics = result.speculative_metrics + rates = ', '.join( + f'{rate:.3f}' + for rate in metrics.per_position_acceptance_rate) + print() + print('--- speculative metrics ---') + print(f'num_drafts: {metrics.num_drafts}') + print(f'num_draft_tokens: {metrics.num_draft_tokens}') + print(f'num_accepted_tokens: {metrics.num_accepted_tokens}') + print( + f'draft_acceptance_rate: ' + f'{metrics.draft_acceptance_rate * 100:.2f}%') + print( + f'mean_acceptance_length: ' + f'{metrics.mean_acceptance_length:.2f}') + print(f'per_position_acceptance_rate: {rates}') print() print('--- tokens ---') for item in result.responses: @@ -611,6 +742,9 @@ def run_smoke_test( cache_prompt: str = DEFAULT_CACHE_PROMPT, cache_generation: str = DEFAULT_CACHE_GENERATION, cache_prompt_boundary_skip: int = DEFAULT_CACHE_PROMPT_BOUNDARY_SKIP, + speculative_method: str | None = None, + speculative_model: str = '', + num_speculative_tokens: int = 1, debug: bool = False, emit_report: bool = True, ) -> SmokeResult: @@ -639,6 +773,9 @@ def run_smoke_test( cache_prompt=cache_prompt, cache_generation=cache_generation, cache_prompt_boundary_skip=cache_prompt_boundary_skip, + speculative_method=speculative_method, + speculative_model=speculative_model, + num_speculative_tokens=num_speculative_tokens, debug=debug, ) if emit_report: @@ -662,6 +799,9 @@ def run_smoke_test( cache_prompt=cache_prompt, cache_generation=cache_generation, cache_prompt_boundary_skip=cache_prompt_boundary_skip, + speculative_method=speculative_method, + speculative_model=speculative_model, + num_speculative_tokens=num_speculative_tokens, debug=debug, ) return result @@ -691,6 +831,9 @@ def main() -> None: cache_prompt=args.cache_prompt, cache_generation=args.cache_generation, cache_prompt_boundary_skip=args.cache_prompt_boundary_skip, + speculative_method=args.speculative_method, + speculative_model=args.speculative_model, + num_speculative_tokens=args.num_speculative_tokens, debug=args.debug, emit_report=True, ) diff --git a/src/turbomind/comm/CMakeLists.txt b/src/turbomind/comm/CMakeLists.txt index c2e320f5e6..102bc327ec 100644 --- a/src/turbomind/comm/CMakeLists.txt +++ b/src/turbomind/comm/CMakeLists.txt @@ -9,6 +9,8 @@ target_link_libraries(host_comm PRIVATE core Threads::Threads) set_property(TARGET host_comm PROPERTY POSITION_INDEPENDENT_CODE ON) add_library(device_comm STATIC device_comm.cc) +target_sources(device_comm PRIVATE + padded_row_allgather.cc) target_link_libraries(device_comm PRIVATE core) set_property(TARGET device_comm PROPERTY POSITION_INDEPENDENT_CODE ON) set_property(TARGET device_comm PROPERTY CUDA_RESOLVE_DEVICE_SYMBOLS ON) diff --git a/src/turbomind/comm/cuda_ipc/allgather.cu b/src/turbomind/comm/cuda_ipc/allgather.cu index 461428523a..d4443cb257 100644 --- a/src/turbomind/comm/cuda_ipc/allgather.cu +++ b/src/turbomind/comm/cuda_ipc/allgather.cu @@ -85,6 +85,11 @@ void CudaIpcCommImpl::AllGather( const int ranks = this->n_ranks(group); const int rank = this->rank(group); + auto* local_slot = static_cast(recvbuff) + rank * bytesize; + if (sendbuff != local_slot && bytesize != 0) { + cudaMemcpyAsync(local_slot, sendbuff, bytesize, cudaMemcpyDeviceToDevice, stream); + } + auto semaphore = groups_.at(group).semaphore.handle(); auto invoke = [&](auto t) { diff --git a/src/turbomind/comm/cuda_ipc/fused_allreduce_ex.cu b/src/turbomind/comm/cuda_ipc/fused_allreduce_ex.cu index 6a534bb5cd..fd8fce9b28 100644 --- a/src/turbomind/comm/cuda_ipc/fused_allreduce_ex.cu +++ b/src/turbomind/comm/cuda_ipc/fused_allreduce_ex.cu @@ -8,6 +8,7 @@ #include "src/turbomind/comm/cuda_ipc/multimem.cuh" +#include "src/turbomind/comm/token_ownership.h" #include "src/turbomind/core/data_type.h" #include "src/turbomind/kernels/core/array_ops.h" #include "src/turbomind/kernels/core/common.h" @@ -320,27 +321,17 @@ void CudaIpcCommImpl::AllreduceResidualBiasRMSnormEx(void* hidden, TM_CHECK(tp0 % inner_tp == 0 && tp1 % inner_tp == 0); - Array offsets{}; - Array firsts{}; - Array lasts{}; + Array ownership{}; - for (int i = 0, offset = 0; i < global_n_ranks_; ++i) { - const int num = local_token_nums[i / inner_tp]; - const int slice = (num + inner_tp - 1) / inner_tp; - const int first = std::min(num, i % inner_tp * slice); - const int last = std::min(num, first + slice); - - std::tie(offsets[i], firsts[i], lasts[i]) = std::tie(offset, first, last); - - if ((i + 1) % inner_tp == 0) { - offset += num; - } + for (int i = 0; i < global_n_ranks_; ++i) { + ownership[i] = ComputeTokenOwnership(i, tp0, tp1, local_token_nums); } const int g_rank = rank(0); - const int first = firsts[g_rank]; - const int last = lasts[g_rank]; - const int offset = offsets[g_rank]; + const auto& owned = ownership[g_rank]; + const int first = owned.local_begin(); + const int last = owned.local_end(); + const int offset = owned.global_offset(); auto semaphore = groups_.at(0).semaphore.handle(); @@ -376,8 +367,9 @@ void CudaIpcCommImpl::AllreduceResidualBiasRMSnormEx(void* hidden, else { Array ag_ranges{}; for (int i = 0; i < tp1; ++i) { - const auto r = g1.l2g[i]; - ag_ranges[i] = {offsets[r] + firsts[r], offsets[r] + lasts[r]}; + const auto r = g1.l2g[i]; + const auto& peer_owned = ownership[r]; + ag_ranges[i] = {peer_owned.global_begin(), peer_owned.global_end()}; } const int max_ctas = max_ctas_.apply(48); AllreduceResidualBiasRMSnormV_Simple_Pull<<>>((T*)hidden, diff --git a/src/turbomind/comm/nccl/nccl.cu b/src/turbomind/comm/nccl/nccl.cu index 278dbfc4d1..310f89a6c7 100644 --- a/src/turbomind/comm/nccl/nccl.cu +++ b/src/turbomind/comm/nccl/nccl.cu @@ -2,7 +2,6 @@ #include #include -#include #include #include @@ -12,6 +11,7 @@ #include "src/turbomind/comm/device_comm.h" #include "src/turbomind/comm/host_comm.h" +#include "src/turbomind/comm/token_ownership.h" #include "src/turbomind/core/check.h" #include "src/turbomind/core/logger.h" #include "src/turbomind/utils/cuda_utils.h" @@ -421,50 +421,36 @@ public: NCCLCHECK(ncclCommCount(comm0, &tp0)); NCCLCHECK(ncclCommCount(comm1, &tp1)); - const int inner_tp = global_n_ranks_ / local_token_nums_count; + const int inner_tp = std::min(tp0, tp1); - std::vector> tasks; + TM_CHECK(tp0 % inner_tp == 0 && tp1 % inner_tp == 0); + + std::vector tasks; tasks.reserve(global_n_ranks_); - for (int i = 0, offset = 0; i < global_n_ranks_; ++i) { - const int num = local_token_nums[i / inner_tp]; - const int slice = (num + inner_tp - 1) / inner_tp; - const int first = std::min(num, i % inner_tp * slice); - const int last = std::min(num, first + slice); - tasks.emplace_back(offset, first, last - first); - if ((i + 1) % inner_tp == 0) { - offset += num; - } + for (int i = 0; i < global_n_ranks_; ++i) { + tasks.push_back(ComputeTokenOwnership(i, tp0, tp1, local_token_nums)); } - const int rank0 = rank(group0); - const int rank1 = rank(group1); - TM_CHECK_EQ(rank0, global_rank_ % tp0); - TM_CHECK_EQ(rank1, global_rank_ % tp1); - - const int rs_begin = global_rank_ - rank0; - const int rs_end = rs_begin + tp0; - - const int ag_begin = global_rank_ - rank1; - const int ag_end = ag_begin + tp1; - // group0: reduce if (tp0 > 1) { NCCLCHECK(ncclGroupStart()); - for (int i = rs_begin; i < rs_end; ++i) { - if (auto& [offset, first, num] = tasks[i]; num > 0) { - char* buff = (char*)hidden + elem_size * (offset + first) * dim; - const int root = i - rs_begin; - NCCLCHECK(ncclReduce(buff, buff, (size_t)num * dim, nccl_type, ncclSum, root, comm0, stream)); + for (int i = 0; i < global_n_ranks_; ++i) { + const auto& owned = tasks[i]; + const int num = owned.row_count(); + if (num > 0) { + char* buff = (char*)hidden + elem_size * owned.global_begin() * dim; + NCCLCHECK(ncclReduce(buff, buff, (size_t)num * dim, nccl_type, ncclSum, i % tp0, comm0, stream)); } } NCCLCHECK(ncclGroupEnd()); } - if (auto& [offset, first, num] = tasks[global_rank_]; num > 0) { - char* buff = (char*)hidden + elem_size * (offset + first) * dim; + const auto& owned = tasks[global_rank_]; + if (const int num = owned.row_count(); num > 0) { + char* buff = (char*)hidden + elem_size * owned.global_begin() * dim; TM_SCOPE_CALL(invokeResidualBiasRMSNorm(buff, - (char*)residual + elem_size * first * dim, + (char*)residual + elem_size * owned.local_begin() * dim, weights, bias, type, @@ -478,11 +464,12 @@ public: // group1: all-gather if (tp1 > 1) { NCCLCHECK(ncclGroupStart()); - for (int i = ag_begin; i < ag_end; ++i) { - if (auto& [offset, first, num] = tasks[i]; num > 0) { - char* buff = (char*)hidden + elem_size * (offset + first) * dim; - const int root = i - ag_begin; - NCCLCHECK(ncclBroadcast(buff, buff, (size_t)num * dim, nccl_type, root, comm1, stream)); + for (int i = 0; i < global_n_ranks_; ++i) { + const auto& peer_owned = tasks[i]; + const int peer_num = peer_owned.row_count(); + if (peer_num > 0) { + char* buff = (char*)hidden + elem_size * peer_owned.global_begin() * dim; + NCCLCHECK(ncclBroadcast(buff, buff, (size_t)peer_num * dim, nccl_type, i % tp1, comm1, stream)); } } NCCLCHECK(ncclGroupEnd()); diff --git a/src/turbomind/comm/padded_row_allgather.cc b/src/turbomind/comm/padded_row_allgather.cc new file mode 100644 index 0000000000..21d6ec3e9d --- /dev/null +++ b/src/turbomind/comm/padded_row_allgather.cc @@ -0,0 +1,52 @@ +#include "src/turbomind/comm/padded_row_allgather.h" + +#include + +#include "src/turbomind/core/data_type.h" +#include "src/turbomind/kernels/core/math.h" + +namespace turbomind { + +Tensor PaddedRowAllGather(Tensor projected_local, + Tensor gathered_padded, + int logical_row_count, + int local_row_count, + int model_tp_rank, + int model_tp_size, + comm::DeviceCommImpl& communicator, + int model_tp_group, + cudaStream_t stream) +{ + const int n = logical_row_count; + const int T = model_tp_size; + const int H = projected_local.shape(1); + + if (n == 0) { + return projected_local.slice({0, 0}, {0, H}); + } + if (T == 1) { + return projected_local.slice({0, 0}, {n, H}); + } + + const int slice = cdiv(n, T); + if (local_row_count < slice) { + cudaMemset2DAsync(static_cast(projected_local.raw_data()) + + local_row_count * projected_local.stride(0) * byte_size(projected_local.dtype()), + projected_local.stride(0) * byte_size(projected_local.dtype()), + 0, + projected_local.shape(1) * byte_size(projected_local.dtype()), + slice - local_row_count, + stream); + } + + communicator.AllGather(projected_local.raw_data(), + gathered_padded.raw_data(), + slice * H, + projected_local.dtype(), + model_tp_group, + stream); + + return gathered_padded.slice({0, 0}, {n, H}); +} + +} // namespace turbomind diff --git a/src/turbomind/comm/padded_row_allgather.h b/src/turbomind/comm/padded_row_allgather.h new file mode 100644 index 0000000000..ce91210cdb --- /dev/null +++ b/src/turbomind/comm/padded_row_allgather.h @@ -0,0 +1,20 @@ +#pragma once + +#include + +#include "src/turbomind/comm/device_comm.h" +#include "src/turbomind/core/core.h" + +namespace turbomind { + +Tensor PaddedRowAllGather(Tensor projected_local, + Tensor gathered_padded, + int logical_row_count, + int local_row_count, + int model_tp_rank, + int model_tp_size, + comm::DeviceCommImpl& communicator, + int model_tp_group, + cudaStream_t stream); + +} // namespace turbomind diff --git a/src/turbomind/comm/token_ownership.h b/src/turbomind/comm/token_ownership.h new file mode 100644 index 0000000000..f6484b4d3f --- /dev/null +++ b/src/turbomind/comm/token_ownership.h @@ -0,0 +1,69 @@ +#pragma once + +#include +#include + +namespace turbomind::comm { + +class OwnedTokenRows { +public: + constexpr OwnedTokenRows() = default; + + constexpr OwnedTokenRows(int global_offset, int local_begin, int local_end): + global_offset_(global_offset), local_begin_(local_begin), local_end_(local_end) + { + } + + constexpr int global_offset() const noexcept + { + return global_offset_; + } + + constexpr int local_begin() const noexcept + { + return local_begin_; + } + + constexpr int local_end() const noexcept + { + return local_end_; + } + + constexpr int row_count() const noexcept + { + return local_end_ - local_begin_; + } + + constexpr int global_begin() const noexcept + { + return global_offset_ + local_begin_; + } + + constexpr int global_end() const noexcept + { + return global_offset_ + local_end_; + } + +private: + int global_offset_{}; + int local_begin_{}; + int local_end_{}; +}; + +inline OwnedTokenRows ComputeTokenOwnership(int global_rank, int tp0, int tp1, const int* local_token_nums) +{ + const int inner_tp = std::min(tp0, tp1); + + const int dp_index = global_rank / inner_tp; + const int tp_index = global_rank % inner_tp; + const int num = local_token_nums[dp_index]; + + const int slice = (num + inner_tp - 1) / inner_tp; + const int first = std::min(num, tp_index * slice); + const int last = std::min(num, first + slice); + const int offset = std::accumulate(local_token_nums, local_token_nums + dp_index, 0); + + return {offset, first, last}; +} + +} // namespace turbomind::comm diff --git a/src/turbomind/engine/CMakeLists.txt b/src/turbomind/engine/CMakeLists.txt index ec6ac903e9..9ae6fff366 100644 --- a/src/turbomind/engine/CMakeLists.txt +++ b/src/turbomind/engine/CMakeLists.txt @@ -3,6 +3,7 @@ cmake_minimum_required(VERSION 3.25) add_library(engine STATIC + model.cc block.cc cache_registry.cc gateway.cc diff --git a/src/turbomind/engine/README.md b/src/turbomind/engine/README.md index 7fae16e964..3d4fc31e21 100644 --- a/src/turbomind/engine/README.md +++ b/src/turbomind/engine/README.md @@ -26,7 +26,18 @@ When code and this document disagree, treat the disagreement as a design bug. Ei ### sequence -`Sequence` is the engine-local mutable execution state for one accepted request on one local rank. It is created from a `Request` during admission and is the object passed through scheduler and model-module contracts. It stores token progress, scheduling decisions, logical block handles, cache-category request state, generation rows, lifecycle flags, and transient per-pass fields. +`Sequence` is the engine-local mutable execution state for one accepted +request on one local rank. It is created from a `Request` during admission and +is the object passed through scheduler and model-module contracts. It stores +token progress, scheduling decisions, logical block handles, cache-category +request state, generation rows, lifecycle flags, and transient per-pass +fields. Its optional `submitted` value is the complete scheduler-to-executor +description of one committed or still-outstanding row. The value contains +input/history length, query and cache-capacity geometry, the `generating` and +`autoregres` execution flags, and the producer-set effect fields +(`verification_positions`, `min_grant`, the inflight completion deltas, +`frontier_reanchor`, and `primes_proposals`). These fields are not stored as +parallel top-level `Sequence` scalars. ### multimodal-spans @@ -41,11 +52,38 @@ When code and this document disagree, treat the disagreement as a design bug. Ei ### phase -Phase is one slot in the async pipeline. With one phase, the engine behaves synchronously: a submitted batch is updated before the next batch is prepared. With multiple phases, host scheduling and setup may run ahead of model execution by reusing different `BatchData` slots. +Phase is one slot in the asynchronous pipeline. With one phase, the engine +behaves synchronously: a submitted batch is updated before the next batch is +prepared. With multiple phases, host scheduling and setup may run ahead of +model execution by reusing different `BatchData` slots. Each phase carries +its reusable `BatchData` and selects module-owned phase buffers; committed +cache-allocation handles are resolved to raw addresses during engine setup +and stored in those per-phase module buffers. `ModelExecutor` consumes +phases in submission order, so every device operation of a phase is +stream-ordered before every operation of its successor; auxiliary-stream +work joins the main stream within its own phase. Changing `CacheBlock` +metadata later cannot change an address already resolved into an earlier +phase's module buffers, deallocating a slot in the preallocated cache region +does not unmap it, and `Engine::Update()` waits for done before the phase +slot is reused. ### scheduler-transaction -Scheduler transaction is one scheduling pass over eligible `Sequence` objects. For each request the scheduler plans (`PlanResume` for inactive, `PlanContinue` for active): it sizes logical blocks, reserves cache block slots, computes `resume_len`, and emits restore copy plans. `Scheduler::Schedule()` then commits: it decides which requests become active, assigns `history_len` and `input_len`, commits cache allocation and eviction through the memory replay, selects and attaches checkpoint publication slots, emits publication copy plans, and records producer marks. +A scheduler transaction is one scheduling pass over eligible `Sequence` +objects. For each request the scheduler plans (`PlanResume` for inactive, +`PlanContinue` for active): it sizes logical blocks, reserves cache block +slots, computes `resume_len`, and emits restore copy plans. +`Scheduler::Schedule()` then commits: it decides which requests become +active, commits cache allocation and eviction through the memory replay, +selects and attaches checkpoint publication slots, emits publication copy +plans, records producer marks, and assigns one complete `SubmittedRow` for +every admitted request. Its input/history lengths, query bounds, cache bounds, +execution flags, and producer-set effect fields (`verification_positions`, +`min_grant`, the inflight completion deltas, `frontier_reanchor`, and +`primes_proposals`; ADR 0003) +are assigned as one value; consumers read effects and never classify rows by +engine mode. An uncommitted request with no outstanding predecessor +has no `submitted` value. ### logical-block @@ -61,7 +99,7 @@ Cache object is an object-typed allocation handle tracked by `CacheBlockPool` an ### module -Module is any TurboMind model component that participates in `LanguageModel::Run(BatchOp, phase, env)`, such as input processing, attention, GDN, generation, or output processing. Modules may validate and prepare their own state, but they must obey the `BatchOp` contracts in this document. +Module is any TurboMind model component that participates in the batch-operation fanouts — the Model's generic fanout and the executor's device-bracket steps — such as input processing, attention, GDN, generation, output processing, or a composed speculative model. Modules may validate and prepare their own state, but they must obey the `BatchOp` contracts in this document. ### signal @@ -73,11 +111,11 @@ Gateway accepts external requests into per-queue `RequestQueue` objects and owns ### engine-thread -Engine thread runs `Engine::Impl::InternalThreadEntry()`. It owns request admission, validation, cancellation observation, scheduling, host-side setup, completed-batch update, lifecycle retirement, and notification submission. All scheduler state is mutated on this thread. +Engine thread runs `Engine::Impl::InternalThreadEntry()`. It owns request admission, validation, cancellation observation, scheduling, the host batch-op functions (`kAdd`, `kSetup`, `kFetch`, `kUpdate`, `kDel`), completed-batch update, lifecycle retirement, and notification submission. All scheduler state is mutated on this thread. ### model-executor-thread -Model executor thread runs `ModelExecutor::Impl::InternalThreadEntry()`. It owns the CUDA execution context for `BatchOp::kPrepare`, `BatchOp::kForward`, and `BatchOp::kUnprep`. It consumes ready `BatchData` objects from the outbound queue, waits for the setup event, runs device work, records the done event, and returns the batch through the inbound queue. +Model executor thread runs `ModelExecutor::Impl::InternalThreadEntry()`. It owns the CUDA execution context for the device bracket's named steps — `BatchOp::kPrepare`, `BatchOp::kForward`, and `BatchOp::kUnprep`. It consumes ready `BatchData` objects from the outbound queue, waits for the setup event, runs device work, records the done event, and returns the batch through the inbound queue. ### data-path @@ -99,11 +137,18 @@ The engine and executor exchange `BatchData` slots through queues. Each slot has ### engine-state -The engine thread is the owner of request scheduling state. It admits requests, mutates `Sequence` lifecycle fields, runs scheduler transactions, calls host-side module-level `BatchOp` handlers, submits batches, processes completed batches, and releases request-owned state. +The engine thread is the owner of request scheduling state. It admits requests, mutates `Sequence` lifecycle fields, runs scheduler transactions, runs the host-side `BatchOp` operations through the Model's generic fanout, submits batches, processes completed batches, and releases request-owned state. ### scheduler-boundary -The scheduler is the transaction boundary for shared execution resources. Request-level planning (`AdmitPrompt`, `PlanResume`, `PlanContinue`) may match or create logical blocks, reserve cache block slots, compute `resume_len`, and emit copy intent, but allocation, eviction, active admission, `history_len`, `input_len`, publication slot attachment, and producer marking are committed by `Scheduler::Schedule()`. +The scheduler is the transaction boundary for shared execution resources and +submitted execution geometry. Request-level planning (`AdmitPrompt`, +`PlanResume`, `PlanContinue`) may match or create logical blocks, reserve +cache block slots, compute `resume_len`, and emit copy intent, but allocation, +eviction, active admission, publication-slot attachment, producer marking, +and the complete `Sequence::submitted` value are committed by +`Scheduler::Schedule()`. Engine code and module setup consume that value; they +do not reconstruct or independently rewrite its fields. ### cache-semantics @@ -115,11 +160,15 @@ Generic cache validity is a lifetime fact, not a resume proof. A valid allocatio ### device-content -Device content operations happen on the model executor thread. Module-specific content work (clearing or post-processing a module's own byte range, preparing pointers, reading model outputs) belongs to the relevant `BatchOp` handler. Whole-object cache copies planned by the scheduler as `(src, dst)` cache-block pairs are resolved to addresses during engine-thread setup and performed by the executor: restore copies before `BatchOp::kPrepare`, publication copies after `BatchOp::kUnprep`. The scheduler never knows what the copied bytes mean; modules never know why a copy happened. Resolving a composite handle yields one or more segments, so a scheduler-planned whole-object copy fans out to one device copy per part (same `(src, dst)` cache-block plan; only the engine-thread resolution multiplies). +Device content operations happen on the model executor thread. Module-specific content work (clearing or post-processing a module's own byte range, preparing pointers, reading model outputs) belongs to the relevant `BatchOp` handler. Whole-object cache copies planned by the scheduler as `(src, dst)` cache-block pairs are resolved to addresses during engine-thread setup and performed by the executor as bracketing steps of its device path: restore copies before the prepare step, publication copies after the unprep step. The scheduler never knows what the copied bytes mean; modules never know why a copy happened. Resolving a composite handle yields one or more segments, so a scheduler-planned whole-object copy fans out to one device copy per part (same `(src, dst)` cache-block plan; only the engine-thread resolution multiplies). ### delayed-cleanup -Async execution requires delayed cleanup. A request that has finished or been canceled must be excluded from future scheduling immediately, but its request-owned resources cannot be released until every submitted batch that references it has completed and decremented `inflight`. +Async execution requires delayed cleanup. A request that has finished or been +canceled is excluded from future scheduling immediately, but its request-owned +resources are released only after every submitted batch that references it +completes and the retiring request reaches `inflight == 0`. Finishing or +canceling the request does not shorten that lifetime. ### callbacks @@ -145,7 +194,16 @@ Partial-block boundary publication is decided entirely at AdmitPrompt-time in `S ### batch-data -`BatchData` slots are owned by the engine/executor queues. A submitted slot temporarily owns the active membership snapshot encoded by `bs0`, `bsz`, and `perm`, plus CUDA events that order setup and execution. It does not own `Sequence` objects. +`BatchData` slots are owned by the engine/executor queues. A submitted slot +temporarily owns the active membership snapshot encoded by `bs0`, `bsz`, and +`perm`, token-count metadata, and CUDA events that order setup and execution. +`BatchData` does not own `Sequence` objects. + +A batch slot retains a handle to the engine-owned symmetric scratch +allocation. Setup copies only this handle. The executor exposes it through +the forward environment after vision processing. The handle preserves +allocation lifetime; scratch contents are shared across phases and reused +in executor order, not owned as persistent per-phase state. ### scheduler @@ -157,11 +215,11 @@ Partial-block boundary publication is decided entirely at AdmitPrompt-time in `S ### module-cache -Modules register anonymous byte requirements with prefix or checkpoint cache categories during construction and keep only byte offsets or base part ids (per registration channel). Each category registers one composite `ObjectAllocator` object id after all modules have registered. A category exposes two registration channels: an accumulation channel (grows part 0, returns a within-part byte offset) and a composite channel (appends parts 1..N, returns the base part id). Slab classes in `ObjectAllocator` are deduped by aligned size, and two same-aligned-size simple categories would share an object id (out of scope: prefix is the only simple category). Modules own the content semantics of their registered byte ranges. The `CacheRegistry` is a registration table only; cache block slot reservation, validity checks, resume selection, and release all live in the scheduler. +Modules register anonymous byte requirements with prefix or checkpoint cache categories during construction and keep only byte offsets or base part ids (per registration channel). Each category registers one composite `ObjectAllocator` object id after all modules have registered. A category exposes two registration channels: an accumulation channel (grows part 0, returns a within-part byte offset) and a composite channel (appends parts 1..N, returns the base part id). Slab classes in `ObjectAllocator` are deduped by aligned size, and two same-aligned-size simple categories would share an object id (out of scope: prefix is the only simple category). Modules own the content semantics of their registered byte ranges. When a speculative model is composed, target attention registers first and draft attention registers second, so they own disjoint byte ranges in the same prefix-category object; this registration order is fixed by construction order in `CreateEngine`. The `CacheRegistry` is a registration table only; cache block slot reservation, validity checks, resume selection, and release all live in the scheduler. ### generation-row -Generation rows are request-owned logical resources managed by the `Generation` module. A row is allocated lazily when a request first generates and is returned only by `BatchOp::kDel` during request cleanup. +Generation rows are request-owned logical resources managed by the `Generation` module. A row is allocated eagerly, before a request's first prompt submission, exactly when the speculative policy reports a bootstrap extent for that prompt length, because the bootstrapping forward writes proposals into that row. A method reporting no bootstrap keeps lazy allocation at first generating submission, as does a target-only engine with no policy. A request with no row yet contributes a null row pointer that every consumer skips. In both modes the row is returned only by `BatchOp::kDel` during request cleanup. ### prefix @@ -187,23 +245,55 @@ Callbacks are owned outside the engine scheduling path. The engine creates signa ### history-len -`history_len` is the committed resume point for the active forward. `Scheduler::Schedule()` sets `history_len = resume_len` only for admitted requests. Module setup and output selection use `history_len` as the start of already-available state for the submitted batch. +`submitted->history_len` is the committed resume point for the submitted +forward. For an ordinary row the scheduler assigns it from `resume_len`. +Module setup and output selection use it as the start of already available +state. It is assigned only for an admitted request; no `submitted` value means +there is no newly committed row. ### input-len -`input_len` is the number of tokens admitted for the active forward. It is set by `Scheduler::Schedule()` after resource admission and allocation planning. Inactive requests must have `input_len == 0` and `history_len == 0`. +`submitted->input_len` is the physical number of tokens admitted for the +submitted forward. For an ordinary row it is assigned after resource and +boundary clamping. An inactive request with no outstanding predecessor has no +`submitted` value. Retaining an outstanding predecessor is not current +admission, and zeroing separate input/history fields is not a state +transition. ### filled-len -`filled_len` is the contiguous prefix context currently established for the request — the position a subsequent resume or decode builds on — not limited to KV this request's own forward produced. It is reconciled in two places. (1) `Engine::Update()` reconciles it from a completed forward: a generating request excludes the newly sampled token, so `filled_len` is `sequence_length - 1`; a non-generating prefill chunk uses `sequence_length`. (2) `Scheduler::CommitResults()` reconciles a resuming request to `filled_len = resume_len`, recording the prefix it reused read-only (prefix cache) or restored from a checkpoint; the `[resume_len, end)` span the in-flight resume forward rebuilds is carried by `inflight_input_len` until that forward completes. The resume-commit write never races `Update()` because a resuming request is inactive (not part of the in-flight batch). +`filled_len` is the contiguous prefix context currently established for the +request, not merely KV produced by that request's own forward. It is +reconciled in two places. `Engine::Update()` uses the completed device +sequence length: a generating row excludes its newly sampled but unconsumed +token and sets `filled_len = returned_sequence_length - 1`, while a +non-generating row sets `filled_len = returned_sequence_length`. +`Scheduler::CommitResults()` reconciles a resuming request to +`filled_len = resume_len`; the in-flight rebuilt span is then represented by +`inflight_input_len` until completion. These writes do not race because a +resuming request is inactive when its resume is committed. Conservative +submitted query/cache bounds and private draft-extension bytes never advance +`filled_len`. ### inflight-input-len -`inflight_input_len` is submitted prefix growth that has not yet been reflected into `filled_len`. In async mode, after update of a completed batch, an active request that was submitted into the next batch records `inflight_input_len = input_len`. This equals `input_len` even for a prefix-skipping resume because `CommitResults()` reconciles `filled_len` to `resume_len`, so the growth the forward produces (`end - filled_len`) is exactly `input_len`. +`inflight_input_len` records submitted prefix growth not yet reflected in +`filled_len`. In async mode, after update of a completed batch, an active +request submitted into the next batch records the submitted row's +`inflight_input_delta`: the full physical width for an ordinary row — still +the full width for a prefix-skipping resume because `CommitResults()` first +reconciles `filled_len = resume_len` — and zero for a speculative row, whose +accepted growth is unknown until device verification completes. Physical +target width and cache capacity remain in `SubmittedRow`. ### inflight-new-tokens -`inflight_new_tokens` is submitted sequence-length growth that has not yet been reflected into `seq_len`. In async mode, after update of a completed batch, an active generating request records `inflight_new_tokens = 1`; otherwise it records `inflight_new_tokens = 0`. +`inflight_new_tokens` records host-predicted sequence growth not yet reflected +in `seq_len`. In async mode, an active ordinary generating successor records +one and every other ordinary successor records zero; the recorded value is the +submitted row's `inflight_new_delta`. A speculative row records +zero; `accept_len` is device-produced and reconciled only after fetch. The +scheduler never predicts speculative acceptance through this field. ### executable-context @@ -211,11 +301,23 @@ The executable context length for a scheduling pass is `seq_len + inflight_new_t ### generating -`generating` means the submitted forward reaches the current context boundary and can produce a next token. The engine sets it from `resume_len + inflight_input_len + input_len == seq_len + inflight_new_tokens`. +`submitted->generating` means the row may produce committed output. Ordinary +rows retain the target-only boundary rule +`resume_len + inflight_input_len + submitted->input_len == seq_len + inflight_new_tokens`. +A speculative row is generating by construction even though its physical +query width is `K`; completed growth is the fetched `accept_len`, not a host +scalar attached at submission. Consumers use the committed flag and do not +rederive it after scheduling. ### autoregres -`autoregres` means the submitted forward is an already-active one-token decode that can take its input token from the model's autoregressive output path instead of copying prompt tokens from host memory. +`submitted->autoregres` means target input IDs are gathered from the persistent +device token row after carried predecessor state is visible. For target-only +execution this remains an already-active one-token generating decode, +classified from the predecessor's committed generating state and +`submitted->input_len == 1`. A speculative row also sets it because its +`K` verification IDs are device-resident proposals. It does not mean the +physical query width is one. ### is-active @@ -239,7 +341,7 @@ The cleanup invariant is: ```cpp if (request.retiring && request.inflight == 0) { - Run(BatchOp::kDel, -1, env); + model_.Run(BatchOp::kDel, -1, env); scheduler.Release(request); remove_sequence(); } @@ -253,7 +355,15 @@ The eviction-protection set a request stamps (`involved_blocks`) is exactly what ### scheduler-start -A scheduler transaction starts with a list of eligible, non-retiring `Sequence` objects. The engine resets transient scheduling fields, and asks the scheduler to plan each request (`PlanResume` for inactive, `PlanContinue` for active) before commit. +A scheduler transaction starts with eligible, non-retiring `Sequence` +objects. The engine resets transient per-pass planning fields and asks the +scheduler to run `PlanResume` for inactive requests and `PlanContinue` for +active requests before commit. Planning may inspect a still-outstanding +`submitted` row; the engine does not clear submitted geometry before planning. +When `inflight == 0`, no outstanding row exists and the committed value is +cleared only if the request is rejected by commit cleanup. Otherwise the +complete submitted row remains available until a successor is committed or the +phase drains. ### prefix-prepare @@ -261,43 +371,122 @@ When prefix caching is enabled and the request is trie-eligible, `Scheduler::Adm ### cache-prepare -Request-level planning (`PlanResume` for inactive requests, `PlanContinue` for active ones) runs inside the scheduling pass before admission. It may create missing logical blocks, reserve missing category cache block slots, compute `resume_len`, and emit restore copy intent as `CacheBlock*` pairs. It must not allocate or deallocate backing object memory, run module callbacks, copy, clear, restore, publish, mark a request active, set `history_len`, or set `input_len`. - -`PlanContinue` maintains the request's `involved_blocks` incrementally rather than rebuilding it: a request active last pass committed, so none of its involved blocks were evicted and its whole required set was allocated; only blocks appended by `EnsureBlocks` since the last plan are new (and, being freshly created, unallocated). `PlanResume` cannot — shared prefix nodes it references can be evicted by other requests between its passes — so it rebuilds `involved_blocks` from a full scan each pass. PlanResume may select an interior partial sibling's checkpoint as the resume point: when the sibling's end lies inside the contiguous valid prefix, its KV range is already covered by the valid full blocks, so planning emits a checkpoint restore copy only (no KV copy); a sibling extending past the prefix end keeps the fork-extension semantics (KV copy plus checkpoint restore when the model checkpoints). +Request-level planning (`PlanResume` for inactive requests, `PlanContinue` for +active ones) runs inside the scheduling pass before admission. It may create +missing logical blocks, reserve missing category cache-block slots, compute +`resume_len`, and emit restore-copy intent as `CacheBlock*` pairs. It does not +commit a `SubmittedRow`, allocate or deallocate backing object memory, run +module callbacks, copy, clear, restore, publish, or mark a request active. + +`PlanContinue` maintains `involved_blocks` incrementally: a request active in +the prior pass committed its required allocation set, so only blocks appended +by `EnsureBlocks` since that plan are new. `PlanResume` rebuilds +`involved_blocks` from a full scan because shared prefix nodes can be evicted +between passes. `PlanResume` may select an interior partial sibling checkpoint +when its end lies inside the contiguous valid prefix; that emits only a +checkpoint restore. A sibling extending past the prefix end retains the +fork-extension behavior of a KV copy plus checkpoint restore when the model +checkpoints. `PlanRequests()` prepares the next row, and required admission +later validates, allocates, and commits it. ### scheduler-commit -`Scheduler::Schedule()` is the commit step. It sorts candidate requests by `Request::unique_id`, stamps each request's `involved_blocks` and the sources of its `restore_copies`, tests composed resources, clamps each forward's end to a boundary candidate (a block boundary, or exactly B = prompt_len - cache_prompt_boundary_skip when `prompt_boundary_node` is set (the publish decision is finalized in `SetupPartialSiblings`; the clamp fires on the pass that can reach `B`); when checkpoint bytes are registered and a prompt-region forward would run past the checkpoint-due position `last_ckpt_pos + checkpoint_min_interval`, its end is truncated to the last block boundary in the admitted range — at or past the due position and strictly past the forward begin — so the full-block checkpoint can be taken there, with the remainder running in the next pass), checks producer conflicts, selects checkpoint publication targets, and plans cache allocation and eviction with a `ScratchAllocator`. Admission and replay run in two phases (see `contracts.scheduler-admission`): `ReplayMemory` is applied once for the required tier and again for the optional tier, and each call applies only its phase's committed replay to the real allocator and then clears the replay buffer. After replay it attaches committed publication slots, emits publication copy plans, updates frontier metadata, and publishes produced ranges. - -For each committed request, the scheduler sets: - -```cpp -r.history_len = r.resume_len; -r.input_len = admitted; // clamped to a boundary candidate: a block boundary, or B = prompt_len - cache_prompt_boundary_skip when prompt_boundary_node is set -r.is_active = true; -``` +`Scheduler::Schedule()` is the commit step. It sorts candidate requests by +`Request::unique_id`, stamps each request's `involved_blocks` and every +restore-copy source, tests +composed resources, and clamps ordinary forward ends to a boundary candidate: +a block boundary, or exactly +`B = prompt_len - cache_prompt_boundary_skip` when +`prompt_boundary_node` is set and the pass can reach `B`. When checkpoint +bytes are registered and a prompt-region forward would run past +`last_ckpt_pos + checkpoint_min_interval`, its end is truncated to the last +block boundary in the admitted range that is at or past the due position and +strictly past the forward begin, so the remainder runs in the next pass. The +commit checks producer conflicts, selects checkpoint publication targets, and +plans cache allocation and eviction with a `ScratchAllocator`. + +Admission and replay retain the two phases from +`contracts.scheduler-admission`. `ReplayMemory` is applied once for the +required tier and again for the optional tier; each application commits only +that tier's replay to the real allocator and clears the replay buffer. After +replay, the scheduler attaches publication slots, emits publication copies, +updates frontier metadata, and publishes produced ranges. + +For every checkpointed submitted row, commit advances live frontier metadata to +the scheduled forward end. A speculative row's K-wide end is a conservative +pending marker, not a reusable exact frontier, and speculative rows are not +checkpoint-publication targets. + +`PlanRequests()` prepares one maximum `SubmittedRow` per request in the existing +`ScheduleState::candidates` storage. It reads the outstanding `submitted` row, +when `inflight != 0`, only to derive the next query and cache-write offsets. It +also applies policy extent and bootstrap geometry before required admission. + +`RunRequiredAdmission()` is uniform. It tests the prepared row once, allows a +positive smaller result to clamp an ordinary prefill, and commits the shortened +row. Speculative resources preserve the zero-or-full-count contract, so a +speculative row cannot enter the partial path. The pass then performs producer +validation, block extension, allocation, eviction, and commit without +re-classifying the row by engine mode or calling the speculative policy. + +`Sequence::submitted = candidate` and `Sequence::is_active = true` occur only +after every required operation succeeds. A speculative candidate is admitted at +exactly its policy query count or not at all. There are no fallback retries: a +resource, producer, allocation, or replay failure leaves the candidate +uncommitted and follows the required-tier failure rule. ### scheduler-inactive -For each uncommitted request, the scheduler must leave it inactive for the current pass: +An uncommitted request is inactive for the current pass. The scheduler clears +its publication target and per-pass allocation/restore/publication-copy +intent. Required admission never overwrites an outstanding `submitted` row +before commit, so cleanup does not restore a captured row. With `inflight == 0`, +cleanup resets `submitted`. Producer conflict may continue to later requests +because it occurs before allocation and replay mutation; `CommitResults()` +clears the uncommitted request's transient vectors. Request-owned logical slots +may remain and are rebuilt by `PlanResume()` on the next pass: ```cpp r.is_active = false; -r.input_len = 0; -r.history_len = 0; r.publish_target = nullptr; r.alloc_blocks.clear(); r.restore_copies.clear(); r.publish_copies.clear(); +if (r.inflight == 0) { + r.submitted.reset(); +} ``` ### scheduler-admission -Admission is two-phase. The **required** tier (prefix blocks + frontier) evicts up to the request's `cutoff[i]` and, on failure, defers the request and stops the pass — priority enforcement, gated by `max_evict_ts`. The **optional** tier (checkpoint publication, fork-to population) runs only after every required forward is placed, on a `ScratchAllocator` (a copy of the committed slab capacity, `MemoryState`; a committed handle is itself the `Allocation` pointer, read for its slot lists during eviction, and the handle store is never copied — `ObjectAllocator` is move-only), and reclaims only **inactive** slots (`timestamp < pass_floor`, where `pass_floor` is the pass-start timestamp) of any category via the allocator's evict/allocate path. An optional allocation that does not fit is dropped; it never evicts active state and never defers a forward. +Admission remains two-phase. The required tier covers prefix blocks plus the +frontier, evicts only up to the request's `cutoff[i]` under `max_evict_ts`, and +on failure defers that request and stops the pass so a lower-priority request +cannot pass it. The optional tier covers checkpoint publication and fork-to +population only after every required forward is placed. It operates on a +`ScratchAllocator`, which copies the committed slab-capacity `MemoryState`. +The committed handle is the `Allocation` pointer itself, read for its slot +lists during eviction; the handle store is never copied and `ObjectAllocator` +remains move-only. The optional tier may reclaim only inactive slots whose +timestamp precedes `pass_floor`. Optional failure drops that optional +allocation; it does not evict active state or defer a required forward. + +`PlanRequests()` classifies the still-outstanding submitted row and prepares one +candidate per request before required admission. The required pass is uniform: +it does not select policy geometry, predict speculative acceptance, or retry a +failed candidate. Ordinary prefill may be shortened when a composed resource +returns a positive partial count. Speculative resources must return zero or the +complete policy query count, so speculative admission is indivisible. A failed +speculative resource or allocation check defers the request and stops the +required tier; it never falls back to an ordinary candidate. ### allocation -Allocation planning must be atomic at the transaction boundary. If a request cannot allocate all required cache objects, the scheduler must not partially mutate the real allocator for that failed suffix. Evictions and allocations are applied only for the committed prefix of the planning replay. +Allocation planning must be atomic at the transaction boundary. If a request +cannot allocate all required cache objects, the scheduler removes only that +request's entries from `pass.planned` and trims the failed replay suffix at +phase cleanup. The real allocator is still mutated only for the committed +prefix when `ReplayMemory()` runs. ### eviction @@ -309,11 +498,41 @@ The scheduler may skip a request whose produced range carries a foreign producer ### scheduler-output -The scheduler's output is a set of current active requests plus updated scheduler metadata. The engine owns batch partitioning, permutation construction, setup submission, update processing, and retirement after the scheduler transaction. +The scheduler outputs current active requests with complete `SubmittedRow` +values plus updated scheduler metadata. The engine owns batch partitioning, +permutation construction, setup submission, update processing, and retirement +after the transaction. A composed speculative engine orders active target +rows as an extension-candidate prefix — speculative rows first, then bootstrap +final prompt rows (possibly multi-token), so the speculative rows form a +leading run — then ordinary generating rows, then ordinary partial-prefill +rows. Extension candidacy is every speculative row plus rows whose forward +primes the first proposals (`primes_proposals`, the scheduler-recorded +bootstrap fold); decoder row partitions rely on the +speculative leading run (ADR 0003). The partition changes +executor row order only; it does not rewrite `SubmittedRow` geometry, +scheduler priority, request ownership, or cache-block order. Target-only +execution retains generating-first order. ### batchop -`BatchOp` is the module-level operation protocol used by `LanguageModel::Run()`. Each operation has a narrow contract. A module may ignore operations that do not apply to it. There is no module-level scheduling operation; scheduler cache preparation owns host-side cache reservation and resume selection. +`BatchOp` is the module-level operation protocol. Each operation's home implies its thread: `kAdd`, `kSetup`, `kFetch`, `kUpdate`, and `kDel` are host operations on the engine thread, while `kPrepare`, `kForward`, and `kUnprep` are the named steps of the executor's device bracket on the model executor thread; none may be ignored by the side that owns it. Each operation has a narrow contract. A module may ignore operations that do not apply to it. There is no module-level scheduling operation; scheduler cache preparation owns host-side cache reservation and resume selection. An unrecognised `spec_method` is a construction-time failure; it is never admitted and then diagnosed during a later operation. + +For every operation the modules are driven in one canonical order: optional +vision, batch status, generation, input processing, target model, optional speculative model, and +output processing. The Model's run method is the generic fanout +for that order — used by every host operation and by the executor's `kUnprep` +step — skipping absent optional modules. `kPrepare` applies the same order +executor-side as a hand-written step carrying its injections: in a composed +engine the verification component's draft inputs are published after +generation's prepare, and the executor publishes the target `k_offsets` buffer +after input preparation and before target-model preparation. Components ignore +operations that do not +apply to them. `kForward` remains an explicit dataflow and +does not use the generic fanout. It is the executor's forward step, branched once per engine composition: +an ordinary engine executes the target-pass routine, a composed engine the +speculative-round routine, which owns the mixed batch (ADR 0002). + +When a speculative model is composed, `output_logits`, `output_last_hidden_state`, `output_logprobs`, `return_ppl`, and guided decoding are unsupported at engine scope rather than conditionally by batch or verification position. Python clears the first four request options. Guided decoding is not cleared or rejected, but the C++ speculative path does not execute it, so asking for it produces output that is not grammar-constrained. ### batchop-add @@ -321,7 +540,22 @@ The scheduler's output is a set of current active requests plus updated schedule ### batchop-setup -`BatchOp::kSetup` runs on the engine thread after scheduler commit and before batch submission. It consumes committed active requests and scheduler metadata. It prepares host and device metadata buffers, copies non-cache input metadata, may resolve committed cache allocation handles to raw addresses, and may update request-owned module handles that describe the submitted work. It must treat the scheduler decision as fixed. +`BatchOp::kSetup` runs on the engine thread after scheduler commit and before +batch submission. It consumes committed active requests, prepares host and +device metadata, copies non-cache input metadata, resolves committed +cache-allocation handles to raw addresses, and may update request-owned module +handles describing the submitted work. Input length, history length, +query/cache capacity, and execution flags come only from the fixed +`SubmittedRow`. It +treats the scheduler decision as fixed: it does not mutate the submitted +value, read device acceptance, or replace conservative capacity with +predicted progress. + +For a composed speculative model, verification positions are the maximum over +generating rows of the row-carried `verification_positions` field (one for an +ordinary generating row, the policy query-row count for a speculative row). +`BatchStatus` derives the count once at `kSetup` and publishes it for every other +module; no forward-time engine module reads `spec_num_draft_tokens`. ### object-address @@ -329,23 +563,152 @@ Resolving an `ObjectAllocator` allocation handle to an address is metadata prepa ### batchop-prepare -`BatchOp::kPrepare` runs on the model executor thread after the setup event is visible on the executor stream and after scheduler-planned restore copies have been enqueued. It prepares device-side state for forward execution. It may use raw cache object addresses prepared by setup and perform module-owned byte-range content operations, such as clearing state for requests whose forward starts at position 0 (`history_len + inflight_input_len == 0`) or post-processing restored content. +`BatchOp::kPrepare` is the prepare step of the executor's device bracket. It +runs on the model executor thread after the setup event is +visible on the executor stream and after the bracket's restore-copy step has +been enqueued. It prepares device-side state, may use raw cache-object +addresses resolved by setup, and may perform module-owned byte-range work such +as zero-start clearing at +`submitted->history_len + inflight_input_len == 0` or post-processing restored +content. When a speculative model is composed, it also carries predecessor `finished` and +`sequence_length`, and the verification component's draft inputs — the +persistent token-row pointers among them — are published as an explicit +bracket step rather than through generation's fanout. It performs no device-to-host transfer or blocking +host synchronization. + +The executor owns target key-offset production in both compositions: it +publishes the `k_offsets` buffer after input preparation and before the target +model prepares, and fills it at `kForward` — over committed sequence lengths in +an ordinary engine, over the input processor's staged per-request key lengths +in a composed one. The input processor's speculative half owns the composed +query-row layout: its forward-time build step gathers the target's input ids +from the request token rows and stages the key lengths, reading the batch's +published operands (`input_ids`, `q_offsets`, `sequence_length`, `finished`, +the verification bracket's `request_token_ids_ptrs`) plus its own +`target_ids_from_row` staging. A buffer borrowed by the target decoder during +`kPrepare` is published before the target prepares and filled at `kForward`. ### batchop-forward -`BatchOp::kForward` runs on the model executor thread. It executes model computation for the submitted batch, mutates module device state for the active requests, writes sampled output ids when generation is active, and updates device-side finished and sequence-length state. KV cache writes are bounded below by `readonly_block_num * logical_block_size`; positions in read-only leading blocks are read but not re-written (`concepts.cache-geometry`). +`BatchOp::kForward` is the forward step of the executor's device bracket: the +visible composition branch (ADR 0002) that runs the speculative-round routine +for a composed engine and the target-pass routine otherwise. It +runs on the model executor thread. It executes model +computation for the submitted batch, mutates module device state, writes +sampled output IDs for generating rows, and updates device-side finished and +sequence-length state. KV stores remain bounded below by +`readonly_block_num * logical_block_size`; leading read-only positions are +read but not rewritten. For a speculative submission, selected target hidden rows +are position-major `[verification_positions, generating_rows, hidden]`. The +executor evaluates the target LM head once over the flattened leading dimensions, +processes the resulting distributions as one block, and performs accept/reject +decisions in position order in one verifier kernel launch containing one CTA per +submitted request. Stop-span clamping precedes the composed speculative model's +draft pass, and persistent `sequence_length` advances exactly once from final +`accept_len`. +The speculative model reads target residuals only through the tap it supplies to the +target decoder. Outside method-owned state, its only persistent cross-round mutations +are its registered cache range and the request token rows. The target pass's +transient staging runs as executor-driven steps before the target decoder: the +input processor's gathered ids and staged key lengths, the executor's +key-offsets prefix-sum, the verification component's selected-states buffer, +and the tap's arming. +Conservative private cache tails may be written but are +not committed prefix progress. No device-written acceptance, token, terminal, +or length value is read by the host in this operation, and executor forward +performs no device-to-host transfer or stream synchronization. + +When target recurrent state is present, speculative target verification +computes all submitted transitions with canonical final-state stores suppressed, +journals rank-local transition inputs, and commits exactly the terminal-clamped +`accept_len` prefix before the draft pass. The journal, accepted length, +and commit remain device-side and stream-ordered. + +For supported unquantized full-attention layers at CP1, a speculative target +verification partition writes K/V through `ProcessKV_v2` and invokes the +standalone CuTe paged-verification kernel once per layer, plus its independent +split reduction when needed. It does not flatten the prefix. Following +ordinary one-token rows retain the existing decode kernel; ordinary prompt +prefill may run on the auxiliary stream and joins before output projection. +Unsupported verification configurations retain whole-batch flattened prefill. +Draft refresh and extension attention remain method-owned prefill/decode work. + +### target-activation + +Inside the target, behavior selects on workload shape or typed caller +arguments — never on the presence of an environment key and never on engine +composition. Layer-internal selection reads row-effect data planned at +`kSetup`: the recurrent-state layer keys its store-suppressed path off its +verification-row count and consumes the store-suppression mask as data, while +the per-row speculative flag flows to device kernels as data only. +Caller-intent capabilities activate through typed `DecoderInputs` fields: a +set `selected_hidden_buffer` is the request to write selected hidden states +into the caller's buffer, available to any caller in any composition. +Environment keys carry data; producing or consuming one at the wrong moment +is a missing-data failure, not a mode change. + +### target-native-capabilities + +`LanguageModel` carries capabilities that exist to serve multi-position +speculative workloads; they are the target's own contract, not speculative +leakage, each with one producer and one consumer: + +- `CommitAcceptedState(phase, accept_len)` — produced by the executor's + speculative round after stop-span clamping; consumed by the target's + recurrent-state layers, which commit exactly the accepted prefix from their + transition journal. +- `SpeculativeStateJournalBytes(request_count, verification_positions)` — + produced at engine construction when sizing transient verification storage; + consumed by the recurrent-state layers' journal and commit sizing. +- `DecoderInputs.taps` — the hidden-state tap the executor passes into the + decoder (the speculator's one target-pass hook); consumed by the decoder's + per-layer capture. +- `DecoderInputs.selected_hidden_buffer` — the caller-owned selected-states + storage, published by the executor from the verification component's buffer + in a composed round; consumed by the decoder's selected-states collection. + `selected_token_pos` is its ordinary-mode sibling. + +The draft-only passthroughs (`attention_input`, `attention_metadata`, and the +`decoder_local_token_nums` topology override) remain the accepted price of +reusing the decoder for the draft (ADR 0001). ### batchop-unprep -`BatchOp::kUnprep` runs on the model executor thread after forward execution and before scheduler-planned publication copies are enqueued. It exports device-side results needed by the engine update path into per-phase module buffers and is the module's last chance to finalize frontier contents before publication snapshots them. It must not invoke external request callbacks. +`BatchOp::kUnprep` is the unprep step of the executor's device bracket, driven +through the Model's generic fanout. It +runs on the model executor thread after the forward step and before +the bracket's publish-copy step. It is the module's last chance to +finalize frontier contents before publication snapshots them and must not +invoke external callbacks. The speculative result export — final phase-owned +selected spans and accepted lengths — happens at `kFetch`, where the +verification component stages them into host-visible buffers for the engine +thread; `kUnprep` itself performs no speculative export. It +is the last module operation before the done event. ### batchop-fetch -`BatchOp::kFetch` runs on the engine thread after the completed batch's done event is visible on the engine stream. It schedules copies from per-phase module buffers to host-visible buffers and publishes fetched tensors into `env` for `kUpdate`. +`BatchOp::kFetch` runs on the engine thread after the completed batch's done +event is visible on the engine stream. It schedules copies from per-phase +module buffers into host-visible buffers and publishes fetched tensors into +`env` for update. Speculative result copies use pinned staging. Fetch does not +mutate `Sequence` or publish cache nodes. ### batchop-update -`BatchOp::kUpdate` runs on the engine thread after fetch copies have completed and the engine stream has synchronized. It updates request-local host state from fetched results and module-owned host buffers. It may update generation sampling state and other CPU-side bookkeeping. It must not release request-owned resources. +`BatchOp::kUpdate` runs on the engine thread after fetch copies complete and +the engine stream synchronizes. It updates request-local host state and +module-owned CPU bookkeeping and must not release request-owned resources. +For a speculative row it maps the completed phase row through its permutation, +reconciles exact `filled_len`, appends exactly `accept_len` committed tokens +for a generating non-retiring request, and finalizes immediately from the +resulting exact prefix. + +Update identifies the completed row from phase-owned `BatchData` — which +snapshots the submitted row's `frontier_reanchor` effect at setup — never from +a possibly overwritten `Sequence::submitted`. After reconciling a completed +checkpointed speculative row, update assigns `frontier_pos = filled_len` only +when no newer phase for that request remains in flight; otherwise it retains the +newer phase's conservative scheduled marker. ### batchop-del @@ -353,7 +716,7 @@ Resolving an `ObjectAllocator` allocation handle to an address is metadata prepa ### executor-only -Only `kPrepare`, `kForward`, and `kUnprep` are executed by the model executor thread. Cache object backing memory is accessed only by these module-level operations and by the executor-run, scheduler-planned whole-object copies that bracket them. +The model executor is the device pipeline only: its public interface is construction and start, built around the engine-owned slot queues. Only the device bracket's `kPrepare`, `kForward`, and `kUnprep` steps are executed by the model executor thread; the host operations run engine-side on the engine thread (ADR 0004). Cache object backing memory is accessed only by these device steps and by the executor-run, scheduler-planned whole-object copies that bracket them. ### cache-metadata @@ -361,7 +724,11 @@ Cache metadata is generic. `CacheBlockPool` owns `CacheBlock` slot storage (stab ### cache-content -Cache contents are module-specific within registered byte ranges. `UnifiedAttentionLayer` owns KV byte-range semantics. `GatedDeltaNetLayer` owns recurrent and convolution state byte-range semantics. Future modules that register category bytes must define their own resumability and content-update rules. +Cache contents are module-specific within registered byte ranges. `UnifiedAttentionLayer` owns KV byte-range semantics. Target and draft KV occupy disjoint ranges in one prefix object, so whole-object allocation, copies, and eviction preserve both ranges together. `GatedDeltaNetLayer` owns recurrent and convolution state byte-range semantics. Future modules that register category bytes must define their own resumability and content-update rules. + +A GDN speculative transition journal is transient executor storage, not cache +content. Only the exact accepted convolution and recurrent state is written to +the checkpoint-category frontier. ### cache-reuse @@ -377,17 +744,124 @@ Modules register anonymous byte requirements with the prefix or checkpoint categ ### unified-attention -`UnifiedAttentionLayer` registers its KV byte requirement with the prefix category during construction and stores the returned byte offset. During setup it resolves committed prefix cache blocks from logical blocks and prepares KV pointer metadata. Reserving logical-block cache slots and validating contiguous prefix coverage is scheduler planning, not module work. Physical KV layout and iteration use `cache_block_seq_len`, while pointer counts and read-only store boundaries use `logical_block_size`. It skips KV cache stores for positions in read-only leading blocks (`< readonly_block_num * logical_block_size`) and supplies those positions from the already-valid blocks during reads (`concepts.cache-geometry`). +`UnifiedAttentionLayer` registers its KV byte requirement with the prefix category during construction and stores the returned byte offset. Target and draft decoders resolve only their own registered byte offsets while sharing the same logical prefix cache block. During setup each layer resolves committed prefix cache blocks from logical blocks and prepares KV pointer metadata. Reserving logical-block cache slots and validating contiguous prefix coverage is scheduler planning, not module work. Physical KV layout and iteration use `cache_block_seq_len`, while pointer counts and read-only store boundaries use `logical_block_size`. It skips KV cache stores for positions in read-only leading blocks (`< readonly_block_num * logical_block_size`) and supplies those positions from the already-valid blocks during reads (`concepts.cache-geometry`). -### gated-deltanet +`ProcessKV_v2` remains the sole multi-query K/V transformation and store +owner. The standalone verification kernel consumes the same `block::Layout` +paged bytes as decode, applies query bias/RoPE/log-N and a position-specific +causal/window mask, and writes either packed output or executor-owned split +partials. It does not participate in the legacy attention registry. +Verification split indices are local to the verification query partition. -`GatedDeltaNetLayer` does not partition its token recurrence across CP ranks. It folds attention CP into GDN tensor parallelism: rank-local weights, convolution state, and recurrent state use `attn_tp_size * attn_cp_size`, with shards selected by `model_tp_rank`. Scheduler checkpoint operations publish and restore these rank-local shards in lockstep at the same global position. +### gated-deltanet -`GatedDeltaNetLayer` registers its recurrent/convolution state byte requirement with the checkpoint category during construction and stores the relevant offsets (per-layer conv element offsets within part 0, computed by the module; the base part id rec_base for recurrent parts). During setup it resolves the committed frontier cache part bases for each request and records which requests start their forward at position 0 (`history_len + inflight_input_len == 0`; in-flight tokens advance the frontier before this batch runs). During `kPrepare` it clears its registered parts (conv part 0 and each recurrent block part, including any rounding padding) for those requests. It does not know whether checkpoints are restored, published, or shared; those are scheduler-planned, executor-run whole-object copies. The recurrent state is a rounded-up 2D `(L_b layers × H_b v_heads)` block grid: one uniform composite part (`block_bytes_`) per block, conv unchanged. `GatedDeltaNetLayer` resolves a per-(layer-group, batch, head-group) recurrent base (composite part `rec_base + (L/L_b)*ng + (h/H_b)`, shared by all `L_b` layers of the block-row) plus a per-layer in-block element offset `linear_state_offset == (L%L_b)*H_b*cell_elems`, and one accumulated conv base (part 0) with the per-layer conv element offset, instead of one recurrent base per layer. The recurrent kernel indexes head-groups: `state_ptrs[b*ng + h/H_b] + linear_state_offset + (h%H_b)*state_size`. With `TM_GDN_BLOCK_CONFIG` unset (`L_b=1, H_b=num_v_heads, ng=1`) this reduces exactly to one base per layer at offset 0. Consumers that reuse a prompt-boundary checkpoint resume at `B` with a restored checkpoint (not position 0), so the "clear at start" path (`history_len + inflight_input_len == 0`) is unaffected. +`GatedDeltaNetLayer` does not partition its token recurrence across CP ranks. +It folds attention CP into GDN tensor parallelism: rank-local weights, +convolution state, and recurrent state use +`attn_tp_size * attn_cp_size`, with shards selected by `model_tp_rank`. +Scheduler checkpoint operations publish and restore these rank-local shards +in lockstep at the same global position. + +`GatedDeltaNetLayer` registers its recurrent/convolution state byte requirement +with the checkpoint category during construction and stores the relevant +offsets (per-layer convolution element offsets within part 0, computed by the +module, and the base part id `rec_base` for recurrent parts). During setup it +resolves the committed frontier cache part bases for each request, reads +`submitted->input_len` as that row's physical input length, and records which +requests start their forward at position zero +(`submitted->history_len + inflight_input_len == 0`; in-flight tokens advance +the frontier before this batch runs). During `kPrepare` it clears its +registered parts (convolution part 0 and each recurrent block part, including +any rounding padding) for those requests. It does not know whether checkpoints +are restored, published, or shared; those are scheduler-planned, executor-run +whole-object copies. + +A speculative target block suppresses the ordinary final convolution and +recurrent-state stores without changing GDN outputs. After target verification +and terminal clamping, the module replays exactly the accepted transition prefix +into its canonical rank-local state and discards the phase journal. + +On eligible SM90 verification inputs of at most 16 positions, the speculative +prefix uses the smallest fitting capacity-8 or capacity-16 single-chunk GDR +forward. It reads the entry recurrent state and emits output without exposing +any state-write path. Other rows retain the ordinary recurrent/chunked kernels. +Transition capture still precedes the forward, and only accepted-prefix replay +mutates canonical convolution and recurrent state after target verification. + +For eligible SM90 inputs of at most 16 positions, recurrent-state commit uses +the same GDR template with layers and speculative requests folded into one +batch. Setup prepares phase-owned host pointers adjusted to each layer's +state slice; prepare copies these pointers and builds their TMA descriptors +on the executor stream. After terminal clamping, final device `accept_len` +is broadcast across layers and the commit kernel updates recurrent state +without producing GDR output. The forward suppression mask and newly updated +`finished` mask do not apply to commit: a newly terminal row still commits +its accepted prefix, while a previously finished row has zero accepted length. +Convolution commit retains its existing ring update, and unsupported recurrent +commit configurations retain scalar replay. Device commit metadata is transient +executor storage, reserved alongside the journal and released on the same +stream after its consumers are enqueued. + +The recurrent state is a rounded-up two-dimensional +`(L_b layers × H_b v_heads)` block grid: one uniform composite part +(`block_bytes_`) per block, with convolution state unchanged. +`GatedDeltaNetLayer` resolves a per-(layer-group, batch, head-group) recurrent +base (composite part +`rec_base + (L/L_b)*ng + (h/H_b)`, shared by all `L_b` layers of the +block-row), plus a per-layer in-block element offset +`linear_state_offset == (L%L_b)*H_b*cell_elems`, and one accumulated +convolution base (part 0) with the per-layer convolution element offset, +instead of one recurrent base per layer. The recurrent kernel indexes +head-groups as +`state_ptrs[b*ng + h/H_b] + linear_state_offset + (h%H_b)*state_size`. +With `TM_GDN_BLOCK_CONFIG` unset +(`L_b=1`, `H_b=num_v_heads`, `ng=1`), this reduces exactly to one base per +layer at offset zero. Consumers that reuse a prompt-boundary checkpoint resume +at `B` with a restored checkpoint, not position zero, so the clear-at-start +path is unaffected. ### checkpoint-publish -Checkpoint publication is planned and committed entirely by the scheduler. Publication targets the node's own block-owned checkpoint slot, created lazily at first publication planning (owner attached at `Create`) and re-allocated in place thereafter — the same model as the prefix slot; no request-owned publication slot exists. Commit knows the forward end only after admitted `input_len`, and planning skips nodes that already hold a validly allocated checkpoint. At most one request can plan a given node per pass — a block target is producer-excluded (the committed forward writes the token before the block end inside it), and a sibling target is reachable only by the request whose trie insert created the boundary node (first-wins arming of `prompt_boundary_node`) — enforced by a checked per-pass reservation of the slot at plan time. Publication planning is routed mutually exclusively by the pass's forward end — a prompt-boundary group (a partial sibling node (`LogicalBlock::partial`)'s KV copy only when `B` is mid-block, plus the boundary checkpoint published either onto that partial sibling or onto the block-aligned boundary block, planned only when `prompt_boundary_node` is set and the forward landed at `B`, so a not-yet-reached pass allocates neither the KV block nor the checkpoint slot) and a full-block group. The full-block group is coverage-driven: it publishes iff a full block ends exactly at the forward end, subject to the configured minimum interval, with no knowledge of prompt-boundary mode. The admission clamp (`contracts.scheduler-commit`) guarantees the full-block group a block-aligned pass end whenever the minimum interval is due in the prompt region, and `PlanResume` seeds `last_ckpt_pos` from a restored checkpoint's position so spacing is measured from it. Recurrent checkpoint publication is suppressed while `is_warm_up` is set (GEMM warm-up); frontier working state is still allocated and updated. Its one exception is `cache_generation=none`, which skips generation-region full blocks (block end `> prompt_len`) while keeping prompt-region full-block checkpoints. The prompt-boundary checkpoint bypasses the minimum interval. Terminal adoption (`contracts.checkpoint-adoption`) may also undercut the interval; the adopted checkpoint is demoted to evict-first priority instead of suppressed. (This drops the prior behavior of suppressing a full-block checkpoint just below an upcoming prompt boundary; that checkpoint is now kept, since full-block publication depends only on coverage.) The optional admission phase allocates the target's checkpoint slot (setting the slot's pin to retain its owner, uniformly with every other owner-attached allocation), and commit records the publication position and emits a frontier-to-slot publication copy that the executor runs after `kUnprep`. +Checkpoint publication is planned and committed entirely by the scheduler. +Publication targets the node's own block-owned checkpoint slot, created lazily +at first publication planning (owner attached at `Create`) and reallocated in +place thereafter, which is the same model as the prefix slot; no request-owned +publication slot exists. Commit knows the forward end only after admitted +`submitted->input_len`, and planning skips nodes that already hold a validly +allocated checkpoint. At most one request can plan a given node per pass: a +block target is producer-excluded because the committed forward writes the +token before the block end inside it, while a sibling target is reachable only +by the request whose trie insert created the boundary node through first-wins +arming of `prompt_boundary_node`. A checked per-pass slot reservation enforces +that uniqueness. + +Publication planning is routed mutually exclusively by the pass's forward end +between a prompt-boundary group and a full-block group. The prompt-boundary +group plans a partial sibling node's KV copy only when `B` is mid-block, plus +the boundary checkpoint onto either that partial sibling or the block-aligned +boundary block. It runs only when `prompt_boundary_node` is set and the +forward lands at `B`, so a not-yet-reached pass allocates neither the KV block +nor the checkpoint slot. The full-block group is coverage-driven: it publishes +if and only if a full block ends exactly at the forward end, subject to the +configured minimum interval and without knowledge of prompt-boundary mode. +The admission clamp in `contracts.scheduler-commit` guarantees a +block-aligned pass end whenever the minimum interval is due in the prompt +region, and `PlanResume` seeds `last_ckpt_pos` from a restored checkpoint so +spacing is measured from it. + +Recurrent checkpoint publication is suppressed while `is_warm_up` is set for +GEMM warm-up; frontier working state is still allocated and updated. The one +exception is `cache_generation=none`, which skips generation-region full +blocks (block end greater than `prompt_len`) while keeping prompt-region +full-block checkpoints. The prompt-boundary checkpoint bypasses the minimum +interval. Terminal adoption in `contracts.checkpoint-adoption` may also +undercut the interval; the adopted checkpoint is demoted to evict-first +priority rather than suppressed. This preserves the current coverage-driven +behavior in which a full-block checkpoint just below a future prompt boundary +is retained. The optional admission phase allocates the target checkpoint +slot, setting the slot pin to retain its owner uniformly with every other +owner-attached allocation. Commit records the publication position and emits +the frontier-to-slot publication copy that the executor runs after `kUnprep`. ### prefix-identity @@ -399,7 +873,17 @@ Producer marking is a per-pass exclusion mechanism. `Scheduler::Schedule()` sets ### prefix-publish -Publication of produced ranges happens at scheduler commit, after the memory replay. Indexed nodes become `is_valid` only when the committed forward end fully covers them; private blocks become `is_valid` with their content extent tracked by `filled_len`. Device-side content arrives in submission order, so a consumer batch always executes after the producer batch that committed before it. +Publication of produced ranges happens at scheduler commit after memory +replay. For ordinary execution, indexed nodes become `is_valid` only when the +committed forward end fully covers them, private blocks become valid with +their content extent tracked by `filled_len`, and device content arrives in +submission order so a consumer executes after the producer batch committed +before it. Speculative target, refresh, and extension writes remain +sequence-private and unindexed while speculative phases are active. +`MarkProduced()` clears producer ownership over the conservative submitted +interval but does not advance `filled_len` or assign prefix identity. Normal +finalization indexes only the exact committed prefix; no delayed prompt +insertion or conservative tail publication exists. ### cancel-release @@ -411,7 +895,28 @@ Terminal checkpoint frontier adoption happens inside `Scheduler::Finalize()` for ### cache-eviction -Eviction may remove cache objects without module-specific knowledge. After eviction, a prefix node remains indexed only while its reference count is positive (requests, fork edges, or remaining valid allocations). Checkpoint and prefix resumability are revalidated by `PlanResume()` on every pass from current allocation validity. Published checkpoints are not held in any request's eviction-protection set (`involved_blocks`), so they age and are reclaimed before live working-set blocks under pressure. While a slot remains demoted (timestamp 0, set by terminal adoption), it sorts before stamped slots and is the first eviction candidate in both admission phases; a later restore/required-use stamp promotes it like any other protected source. Eviction frees a cache allocation, not the `CacheBlock` slot or the `LogicalBlock`: a block referenced by a living sequence or a fork edge survives even with all of its allocations evicted, and is recycled only when its last reference drops. +Eviction may remove cache objects without module-specific knowledge. After +eviction, a prefix node remains indexed only while its reference count is +positive through requests, fork edges, or remaining valid allocations. +`PlanResume()` revalidates checkpoint and prefix resumability on every pass +from current allocation validity. Published checkpoints are not held in a +request's `involved_blocks`, so they age and are reclaimed before live +working-set blocks. A terminal-adopted slot demoted to timestamp zero sorts +before stamped slots until a later restore or required-use stamp promotes it. +Eviction frees a cache allocation, not its `CacheBlock` slot or +`LogicalBlock`; a request or fork reference may keep the block alive after all +allocations are gone, and the block is recycled only when its last reference +drops. + +## non-normative examples + +### eagle3 + +EAGLE3 is one implementation of the speculative seams above, not part of their +contract. For `k` draft tokens, its policy returns `Extent = {k + 1, k - 1}` and +`Bootstrap(prompt_len) = {prompt_len + k - 1, prompt_len + 2 * k}`. Its hidden-state +tap captures residuals at the configured target layer ids. Its draft pass runs one +shifted refresh followed by a serial loop of `k - 1` draft extensions. ## checklist @@ -419,7 +924,7 @@ Before changing TurboMind async execution, scheduler, cache management, or modul ### state-owner -Does exactly one component own each state mutation? +Does exactly one module own each state mutation? ### cache-prepare @@ -427,7 +932,10 @@ Do `AdmitPrompt`/`PlanResume`/`PlanContinue` only match or create logical blocks ### scheduler-commit -Does `Scheduler::Schedule()` remain the only active-admission, allocation, eviction, `history_len`, `input_len`, and publication-attach commit point? +Does `Scheduler::Schedule()` remain the only active-admission, allocation, +eviction, publication-attachment, and `SubmittedRow` commit point? Does every +consumer use the committed value rather than parallel fields or post-scheduler +rederivation? ### cache-semantics @@ -439,7 +947,14 @@ Is generic cache validity used only for lifetime, not to raise `resume_len`? ### cache-memory -Are cache object backing-memory reads and writes limited to executor-thread `BatchOp` handlers and executor-run, scheduler-planned whole-object copies, with KV writes further limited to `[readonly_block_num * logical_block_size, end)` (read-only leading blocks are reads only), while physical KV object sizing and iteration remain based on `cache_block_seq_len` (`concepts.cache-geometry`)? For composite objects, are whole-object copies issued as one device copy per part? +Are cache-object backing-memory reads and writes limited to executor-thread +`BatchOp` handlers and executor-run, scheduler-planned whole-object copies? +Are KV writes limited to +`[readonly_block_num * logical_block_size, end)`, with physical KV sizing and +iteration still based on `cache_block_seq_len`? Are composite whole-object +copies issued once per part? Are speculative writable +destinations private, keyless, unindexed, and bounded by their committed +`SubmittedRow`? ### delayed-release @@ -447,11 +962,15 @@ Can a finishing or canceled request be excluded from scheduling before its resou ### cleanup -Is every request-owned resource released only after `retiring && inflight == 0`? +Is every request-owned resource released only after +`retiring && inflight == 0`? ### async-progress -Does async state account for submitted-but-not-yet-reflected work through `inflight_input_len`, `inflight_new_tokens`, and `inflight`? +Does async state account for submitted but not yet reflected ordinary work +through `inflight_input_len`, `inflight_new_tokens`, and `inflight`? Does host +accounting avoid predicting speculative acceptance and retain exact predecessor +`SubmittedRow` geometry while a phase is outstanding? ### forward-progress diff --git a/src/turbomind/engine/batch.h b/src/turbomind/engine/batch.h index 1f7de3a42b..1d09c8fd26 100644 --- a/src/turbomind/engine/batch.h +++ b/src/turbomind/engine/batch.h @@ -68,12 +68,17 @@ struct BatchData { Buffer_ perm; + // Shared language scratch; contents are reused in executor order. + Buffer_ symm_buf; + std::vector restore_copies; // run before BatchOp::kPrepare std::vector publish_copies; // run after BatchOp::kUnprep std::vector local_token_num; int global_token_num = 0; + std::vector submitted_frontier_reanchor; + Event ready; Event done; Event next; diff --git a/src/turbomind/engine/engine.cc b/src/turbomind/engine/engine.cc index 833ee9466e..1b213d6a72 100644 --- a/src/turbomind/engine/engine.cc +++ b/src/turbomind/engine/engine.cc @@ -4,6 +4,7 @@ #include #include #include +#include #include #include @@ -15,6 +16,7 @@ #include "src/turbomind/core/check.h" #include "src/turbomind/core/context.h" #include "src/turbomind/engine/engine.h" +#include "src/turbomind/engine/model.h" #include "src/turbomind/engine/model_executor.h" #include "src/turbomind/engine/request.h" #include "src/turbomind/engine/scheduler.h" @@ -22,9 +24,12 @@ #include "src/turbomind/core/copy.h" #include "src/turbomind/core/logger.h" #include "src/turbomind/core/scope.h" +#include "src/turbomind/kernels/sampling_topp_kernels.h" #include "src/turbomind/models/language_model.h" #include "src/turbomind/models/llama/context_token_resource.h" #include "src/turbomind/models/llama/llama_params.h" +#include "src/turbomind/models/model_weight.h" +#include "src/turbomind/models/speculative/speculative_model.h" #include "src/turbomind/models/vision_model.h" #include "src/turbomind/utils/cuda_utils.h" #include "src/turbomind/utils/metrics.h" @@ -61,16 +66,16 @@ struct Engine::Impl { struct State; - Impl(EngineParam param, - ObjectAllocator alloc, - CacheRegistry cache_registry, - LanguageModel model, - std::unique_ptr vision_model, - Context& ctx, - Gateway& gateway, - int device_id, - int queue_id, - int phases); + Impl(EngineParam param, + CacheRegistry cache_registry, + std::unique_ptr model, + std::unique_ptr vision_model, + std::unique_ptr spec_model, + Context& ctx, + Gateway& gateway, + int device_id, + int queue_id, + int phases); void InternalThreadEntry(); @@ -95,19 +100,12 @@ struct Engine::Impl { // Initialize batch data from engine-local sequence state void Setup(BatchData& d); - // Sync vars from batch output to engine-local sequence state + // Sync vars from batch output to engine-local sequence state. Host batch + // operations (add, setup, fetch, update, del) run on this thread through + // the component container's generic fanout; device operations are the + // executor's device-bracket steps. void Update(BatchData& d, std::vector& signals); - void Run(BatchOp op, int phase, Ref env) - { - // Vision sub-graph runs first so its env outputs (image embeddings, - // mrope tensors) are visible to the language model in the same pass. - if (vision_model_) { - vision_model_->Run(op, phase, env); - } - model_.Run(op, phase, env); - } - void Start() { internal_thread_ = std::thread(&Impl::InternalThreadEntry, this); @@ -146,14 +144,20 @@ struct Engine::Impl { int& is_warm_up_; ObjectAllocator object_allocator_; - Scheduler scheduler_; + + Buffer_ symm_buf_; + + // The served model, constructed once here at the composition root. It + // owns the target, vision, and speculative models; the executor drives it + // by reference. + Model model_; + + std::unique_ptr scheduler_; Queue> inbound_; Queue> outbound_; - LanguageModel model_; - std::unique_ptr vision_model_; // null for text-only checkpoints - ModelExecutor executor_; + ModelExecutor executor_; std::thread internal_thread_; @@ -208,23 +212,23 @@ Engine::Impl::~Impl() for (auto& state : states_) { for (auto& cache : state.rc) { if (cache) { - scheduler_.Release(*cache); + scheduler_->Release(*cache); cache.reset(); } } } } -Engine::Impl::Impl(EngineParam param, - ObjectAllocator alloc, - CacheRegistry cache_registry, - LanguageModel model, - std::unique_ptr vision_model, - Context& ctx, - Gateway& gateway, - int device_id, - int queue_id, - int phases): +Engine::Impl::Impl(EngineParam param, + CacheRegistry cache_registry, + std::unique_ptr model, + std::unique_ptr vision_model, + std::unique_ptr spec_model, + Context& ctx, + Gateway& gateway, + int device_id, + int queue_id, + int phases): param_{param}, gateway_{gateway}, tp_group_{ctx.comm.h_tp_group}, @@ -236,25 +240,85 @@ Engine::Impl::Impl(EngineParam param, queue_id_{queue_id}, async_{phases > 1}, is_warm_up_{*ctx.is_warm_up}, - object_allocator_{std::move(alloc)}, - scheduler_{object_allocator_, - std::move(cache_registry), - param_.cache_block_seq_len * param_.attn_cp_size, - param_.enable_prefix_caching, - param_.cache_prompt, - param_.cache_prompt_boundary_skip, - param_.cache_generation, - is_warm_up_}, - model_{std::move(model)}, - vision_model_{std::move(vision_model)} + model_{std::move(model), std::move(vision_model), std::move(spec_model), param_, ctx, phases} { - states_.emplace_back(); + const double cache_ratio = param_.cache_max_block_count; + TM_CHECK_GT(cache_ratio, 0.) << "object-cache path expects 0 < cache_max_block_count < 1"; + TM_CHECK_LT(cache_ratio, 1.) << "object-cache path no longer accepts cache_max_block_count as a block count"; + states_.emplace_back(); for (int i = 0; i < phases; ++i) { data_.emplace_back(); } - executor_ = ModelExecutor{model_, vision_model_.get(), ctx, device_id_, outbound_, inbound_}; + executor_ = ModelExecutor{model_, param_, ctx, device_id_, outbound_, inbound_}; + + const ModelWeight& target_weights = model_.target->weights(); + const int max_verification_positions = param_.spec_method.empty() ? 1 : param_.spec_num_draft_tokens + 1; + const int max_logits_rows = param_.max_batch_size * max_verification_positions; + + if (ctx.comm.d_comm) { + const int model_tp_size = ctx.comm.h_tp_group->n_ranks(); + TM_CHECK(param_.max_forward_token_num % model_tp_size == 0); + + const core::ssize_t bytes = std::max( + byte_size(target_weights.data_type, + core::ssize_t(param_.max_forward_token_num) * param_.attn_dp_size * target_weights.hidden_units), + byte_size(target_weights.data_type, core::ssize_t(max_logits_rows) * target_weights.vocab_size_padded)); + + auto symm_alloc = GetSymmAllocator(ctx.comm.d_comm); + symm_buf_ = {bytes, symm_alloc}; + } + + core::Context::stream().Sync(); + + size_t free_after_workspaces{}, total_bytes{}; + cudaMemGetInfo(&free_after_workspaces, &total_bytes); + free_after_workspaces = AllReduce(ctx.comm.h_tp_group, free_after_workspaces, comm::RedOp::kMin); + + size_t transient_verification_bytes{}; + if (model_.spec) { + const size_t vocab_items = static_cast(max_logits_rows) * target_weights.vocab_size_padded; + const size_t target_head_bytes = + model_.target->logits_use_workspace() ? 0 : byte_size(target_weights.data_type, vocab_items); + + transient_verification_bytes = target_head_bytes + vocab_items * sizeof(float) + vocab_items * sizeof(int) + + GetTopPSortWorkspaceBytes(max_logits_rows, + target_weights.vocab_size, + target_weights.vocab_size_padded, + core::Context::stream().handle()); + transient_verification_bytes += + model_.target->SpeculativeStateJournalBytes(param_.max_batch_size, max_verification_positions); + } + + TM_CHECK_GE(free_after_workspaces, transient_verification_bytes) + << "insufficient free memory for transient verification storage"; + const size_t cacheable_bytes = free_after_workspaces - transient_verification_bytes; + const size_t cache_bytes = static_cast(static_cast(cacheable_bytes) * cache_ratio); + + TM_LOG_INFO("Object cache memory: free after model, components, executor, and shared scratch allocation {:.2f} MB, " + "transient verification reservation {:.2f} MB", + free_after_workspaces / (1024. * 1024.), + transient_verification_bytes / (1024. * 1024.)); + TM_LOG_INFO("Object cache budget: {:.2f} MB from cacheable {:.2f} MB and ratio {:.3f}", + cache_bytes / (1024. * 1024.), + cacheable_bytes / (1024. * 1024.), + cache_ratio); + + Buffer cache_region{static_cast(cache_bytes), data_type_v, core::Context::device_alloc()}; + object_allocator_ = ObjectAllocator{std::move(cache_region)}; + cache_registry.RegisterObjectIds(object_allocator_); + + scheduler_ = std::make_unique(object_allocator_, + std::move(cache_registry), + param_.cache_block_seq_len * param_.attn_cp_size, + param_.enable_prefix_caching, + param_.cache_prompt, + param_.cache_prompt_boundary_skip, + param_.cache_generation, + param_.session_len, + model_.spec ? &model_.spec->policy() : nullptr, + is_warm_up_); UpdateScheduleMetrics(); @@ -327,9 +391,10 @@ void Engine::Impl::Interrupt(Sequence& c) { Sequence* p = &c; Buffer_ rs{&p, 1, kCPU}; - Run(BatchOp::kDel, -1, TensorMap{{"requests", rs}}); + TensorMap env{{"requests", rs}}; + model_.Run(BatchOp::kDel, -1, env); - scheduler_.Release(c); + scheduler_->Release(c); } void Engine::Impl::Retire(State& s) @@ -357,7 +422,7 @@ void Engine::Impl::Cancel(vector& indices, vector& signals) c->is_canceled = true; c->retiring = true; c->done = true; - signals.push_back([r = c->req, l = c->seq_len] { UpdateState(*r, Request::kCancel, l); }); + signals.push_back(MakeRequestSignal(c->req, Request::kCancel, c->seq_len)); } } @@ -371,7 +436,7 @@ void Engine::Impl::Accept(const Requests& rs, vector& signals) for (const auto& r : rs) { if (r->ec) { - signals.push_back([r] { UpdateState(*r, r->ec, 0); }); + signals.push_back(MakeRequestSignal(r, r->ec, 0)); continue; } @@ -379,7 +444,7 @@ void Engine::Impl::Accept(const Requests& rs, vector& signals) const int input_len = input_ids.shape(0); if (input_len > param_.session_len) { - signals.push_back([r] { UpdateState(*r, Request::kTooLong, 0); }); + signals.push_back(MakeRequestSignal(r, Request::kTooLong, 0)); continue; } @@ -415,18 +480,17 @@ void Engine::Impl::Accept(const Requests& rs, vector& signals) } // This includes checks from all modules handling `Add` operation - Run(BatchOp::kAdd, -1, TensorMap{{"requests", buf}}); + TensorMap env{{"requests", buf}}; + model_.Run(BatchOp::kAdd, -1, env); for (auto& x : incoming) { if (x->status == 0) { - scheduler_.AdmitPrompt(*x); + scheduler_->AdmitPrompt(*x); s.rc.push_back(std::move(x)); } else { Interrupt(*x); - signals.push_back([r = x->req, ec = x->status] { // - UpdateState(*r, ec, 0); - }); + signals.push_back(MakeRequestSignal(x->req, x->status, 0)); } } } @@ -439,9 +503,7 @@ void Engine::Impl::Schedule() vector eligible; vector was_active; - vector context_length; vector orignal_idxs; - vector inflight_input_len; for (int i = 0; i < s.size(); ++i) { auto& p = s.rc[i]; @@ -452,18 +514,20 @@ void Engine::Impl::Schedule() if (!c.retiring) { eligible.push_back(&c); was_active.push_back(c.is_active); - context_length.push_back(c.seq_len + c.inflight_new_tokens /* plus draft tokens */); - inflight_input_len.push_back(c.inflight_input_len); orignal_idxs.push_back(i); - c.input_len = c.history_len = 0; } } ScheduleResources resources; resources.Add(param_.max_forward_token_num); - resources.Add(param_.max_context_token_num); + // A speculative row is charged its absolute key_capacity_end, which runs up + // to two proposal windows past the last sequence position; keep the context + // budget above that so near-limit rounds stay admissible. + const int context_headroom = + model_.spec ? 2 * model_.spec->policy().Extent(RoundRequest{0, param_.session_len}).query_rows : 0; + resources.Add(param_.max_context_token_num + context_headroom); - scheduler_.Schedule(eligible, resources); + scheduler_->Schedule(eligible, resources); vector idxs(eligible.size()); std::iota(idxs.begin(), idxs.end(), 0); @@ -475,22 +539,17 @@ void Engine::Impl::Schedule() // FailStalledHeadOfLine, called after Schedule() returns, where request // lifecycle and signal emission live (see README forward-progress). - if (is_warm_up_) { - // Avoid extra iteration for warm up request in async mode (force inactivate) - active = {active.begin(), std::stable_partition(active.begin(), active.end(), [&](int i) { - return inflight_input_len[i] == 0; - })}; - } - subrange inactive{active.end(), idxs.end()}; for (auto i : active) { eligible[i]->is_active = true; } for (auto i : inactive) { - eligible[i]->is_active = false; - eligible[i]->input_len = 0; - eligible[i]->history_len = 0; + Sequence& c = *eligible[i]; + c.is_active = false; + if (c.inflight == 0) { + c.submitted.reset(); + } } subrange existing{active.begin(), @@ -504,11 +563,6 @@ void Engine::Impl::Schedule() // |<-- existing -->|<-- swap-in -->|<- swap-out ->| // |<----------- active ----------->|<------- inactive ----->| - for (auto i : swap_in) { - eligible[i]->autoregres = {}; - eligible[i]->generating = {}; - } - for (auto i : swap_in) { auto& c = *eligible[i]; if (!param_.enable_metrics || c.first_schedule_recorded || !c.req->metrics) { @@ -516,7 +570,7 @@ void Engine::Impl::Schedule() } c.first_schedule_recorded = true; - const int64_t cached_tokens = std::clamp(c.history_len, 0, c.prompt_len); + const int64_t cached_tokens = std::clamp(c.submitted->history_len, 0, c.prompt_len); if (!is_warm_up_ && param_.enable_prefix_caching) { prefix_query_tokens_ += c.prompt_len; prefix_hit_tokens_ += cached_tokens; @@ -528,20 +582,15 @@ void Engine::Impl::Schedule() m.scheduled_time.compare_exchange_strong(expected, RequestMetrics::timestamp(), std::memory_order_relaxed); } - for (auto i : existing) { - auto& c = *eligible[i]; - c.autoregres = c.generating && c.input_len == 1; - } + auto extension_end = std::stable_partition( + active.begin(), active.end(), [&](int i) { return eligible[i]->submitted->is_extension_candidate(); }); - for (auto i : active) { - auto& c = *eligible[i]; - c.generating = c.resume_len + c.inflight_input_len + c.input_len == c.seq_len + c.inflight_new_tokens; - } + // Speculative rows precede bootstrap rows inside the extension prefix; decoder + // partitions rely on the speculative rows forming a leading run. + std::stable_partition( + active.begin(), extension_end, [&](int i) { return eligible[i]->submitted->is_verification_row(); }); - // move partially prefilled sequences to the back - subrange partial{ - std::stable_partition(active.begin(), active.end(), [&](int i) { return eligible[i]->generating; }), - active.end()}; + std::stable_partition(extension_end, active.end(), [&](int i) { return eligible[i]->submitted->generating; }); // dbg(inv); @@ -608,7 +657,7 @@ void Engine::Impl::FailStalledHeadOfLine(std::vector& signals) victim->retiring = true; victim->done = true; - signals.push_back([r = victim->req] { UpdateState(*r, Request::kOutOfMemory, 0); }); + signals.push_back(MakeRequestSignal(victim->req, Request::kOutOfMemory, 0)); } void Engine::Impl::Setup(BatchData& d) @@ -616,8 +665,9 @@ void Engine::Impl::Setup(BatchData& d) TM_FUNCTION_SCOPE(); auto& s = states_.at(0); - d.bs0 = s.bs0; - d.bsz = s.active; + d.bs0 = s.bs0; + d.bsz = s.active; + d.symm_buf = symm_buf_; d.perm = {d.bsz, kCPU}; std::copy_n(s.perm.data(), d.bsz, d.perm.data()); @@ -625,16 +675,18 @@ void Engine::Impl::Setup(BatchData& d) BatchCopy copy{}; Buffer_ rs{s.active, kCPU}; + d.submitted_frontier_reanchor.resize(d.bsz); for (int i = 0; i < s.active; ++i) { auto* c = TM_CHECK_NOTNULL(s.rc[i].get()); ++c->inflight; - rs[i] = c; + rs[i] = c; + d.submitted_frontier_reanchor[i] = c->submitted->frontier_reanchor; } d.restore_copies.clear(); d.publish_copies.clear(); { - const ObjectAllocator& alloc = scheduler_.allocator(); + const ObjectAllocator& alloc = scheduler_->allocator(); auto resolve = [&](std::vector& in, std::vector& out) { for (const auto& [src, dst] : in) { const CacheBlock& cs = *TM_CHECK_NOTNULL(src); @@ -659,7 +711,7 @@ void Engine::Impl::Setup(BatchData& d) TensorMap env{{"requests", rs}, {"batch", d.buf()}, {"copy", copy.buf()}}; - Run(BatchOp::kSetup, d.phase, env); + model_.Run(BatchOp::kSetup, d.phase, env); // dbg(copy); copy.Run(); @@ -682,20 +734,39 @@ void Engine::Impl::Update(BatchData& b, std::vector& signals) TensorMap env{{"batch", b.buf()}, {"copy", copy.buf()}}; // Copy outputs to host buffers - Run(BatchOp::kFetch, b.phase, env); + model_.Run(BatchOp::kFetch, b.phase, env); copy.Run(); core::Context::stream().Sync(); // - Run(BatchOp::kUpdate, b.phase, env); + model_.Run(BatchOp::kUpdate, b.phase, env); Buffer_ finished = env.at("finished").buffer(); Buffer_ generating = env.at("generating").buffer(); - Buffer_ output_ids = env.at("output_ids").buffer(); Buffer_ sequence_length = env.at("sequence_length").buffer(); + Buffer_ output_ids; + Buffer_ selected_span_ids; + Buffer_ accept_len; + + const bool speculative_engine = model_.spec != nullptr; + const int K = speculative_engine ? model_.spec->policy().max_proposals() + 1 : 0; + + if (speculative_engine) { + selected_span_ids = env.at("selected_span_ids").buffer(); + accept_len = env.at("accept_len").buffer(); + } + else { + output_ids = env.at("output_ids").buffer(); + } + + Buffer_ accepted_draft_count; + if (const Tensor* tensor = env.try_("accepted_draft_count")) { + accepted_draft_count = Buffer_{tensor->buffer()}; + } + env = {}; vector perm(s.size()); @@ -709,27 +780,59 @@ void Engine::Impl::Update(BatchData& b, std::vector& signals) for (int i = 0; i < s.size(); ++i) { int j = perm[i]; if (j < b.bsz) { - auto& c = *TM_CHECK_NOTNULL(s.rc[i]); - c.filled_len = generating[j] ? sequence_length[j] - 1 : sequence_length[j]; + auto& c = *TM_CHECK_NOTNULL(s.rc[i]); + c.filled_len = generating[j] ? sequence_length[j] - 1 : sequence_length[j]; + const bool completed_frontier_reanchor = b.submitted_frontier_reanchor[j]; + if (speculative_engine && completed_frontier_reanchor && scheduler_->registry().has_checkpoint() + && c.inflight == 1) { + c.frontier_pos = c.filled_len; + } if (c.retiring) { continue; } if (generating[j]) { - c.token_ids[c.seq_len] = output_ids[j]; - c.seq_len = sequence_length[j]; + if (speculative_engine) { + const int committed = accept_len[j]; + std::copy_n(selected_span_ids.data() + j * K, committed, c.token_ids + c.seq_len); + } + else { + c.token_ids[c.seq_len] = output_ids[j]; + } + + c.seq_len = sequence_length[j]; + if (int new_tokens = c.seq_len - c.tokens.size(); TM_LIKELY(new_tokens)) { c.tokens.insert(c.tokens.end(), c.token_ids + c.seq_len - new_tokens, c.token_ids + c.seq_len); } + + if (accepted_draft_count) { + const int accepted = accepted_draft_count[j]; + const int k = model_.spec->policy().max_proposals(); + + if (accepted >= 0 && param_.model_tp_rank == 0 && c.req->metrics) { + auto& metrics = *c.req->metrics; + std::scoped_lock lock(metrics.spec_mutex); + + ++metrics.num_drafts; + metrics.num_draft_tokens += k; + metrics.num_accepted_tokens += accepted; + + for (int pos = 0; pos < accepted; ++pos) { + ++metrics.num_accepted_tokens_per_pos[pos]; + } + } + } + if (TM_UNLIKELY(finished[j])) { if (!c.is_canceled) { - scheduler_.Finalize(c); + scheduler_->Finalize(c); } - signals.push_back([r = c.req, l = c.seq_len] { UpdateState(*r, Request::kFinish, l); }); + signals.push_back(MakeRequestSignal(c.req, Request::kFinish, c.seq_len)); c.retiring = true; c.done = true; } else if (TM_LIKELY(c.req->stream_output)) { - signals.push_back([r = c.req, l = c.seq_len] { UpdateState(*r, Request::kOk, l); }); + signals.push_back(MakeRequestSignal(c.req, Request::kOk, c.seq_len)); } } } @@ -744,8 +847,9 @@ void Engine::Impl::Update(BatchData& b, std::vector& signals) for (int i = 0; i < size; ++i) { auto& c = *s.rc[i]; if (i < s.active) { - c.inflight_input_len = c.input_len; - c.inflight_new_tokens = c.generating; + const SubmittedRow& row = *c.submitted; + c.inflight_input_len = row.inflight_input_delta; + c.inflight_new_tokens = row.inflight_new_delta; } else { // Just got swaped-out @@ -884,25 +988,25 @@ Engine::Engine() = default; Engine::Engine(Engine&&) noexcept = default; Engine& Engine::operator=(Engine&&) noexcept = default; -Engine::Engine(EngineParam param, - ObjectAllocator alloc, - CacheRegistry cache_registry, - LanguageModel model, - std::unique_ptr vision_model, - Context& ctx, - Gateway& gateway, - int device_id, - int dp_rank, - int phases): +Engine::Engine(EngineParam param, + CacheRegistry cache_registry, + std::unique_ptr model, + std::unique_ptr vision_model, + std::unique_ptr spec_model, + Context& ctx, + Gateway& gateway, + int device_id, + int queue_id, + int phases): impl_{std::make_unique(param, - std::move(alloc), std::move(cache_registry), std::move(model), std::move(vision_model), + std::move(spec_model), ctx, gateway, device_id, - dp_rank, + queue_id, phases)} { } diff --git a/src/turbomind/engine/engine.h b/src/turbomind/engine/engine.h index cf162b838c..ff082b0e78 100644 --- a/src/turbomind/engine/engine.h +++ b/src/turbomind/engine/engine.h @@ -13,6 +13,7 @@ namespace turbomind { struct ScheduleMetrics; +class SpeculativeModel; class VisionModel; class Engine { @@ -28,16 +29,16 @@ class Engine { return static_cast(impl_); } - Engine(EngineParam param, - ObjectAllocator alloc, - CacheRegistry cache_registry, - LanguageModel model, - std::unique_ptr vision_model, // null for text-only checkpoints - Context& ctx, - Gateway& gateway, - int device_id, - int queue_id, - int phases); + Engine(EngineParam param, + CacheRegistry cache_registry, + std::unique_ptr model, + std::unique_ptr vision_model, + std::unique_ptr spec_model, + Context& ctx, + Gateway& gateway, + int device_id, + int queue_id, + int phases); void Start(); diff --git a/src/turbomind/engine/engine_config.h b/src/turbomind/engine/engine_config.h index a17af1c7da..725ca2f073 100644 --- a/src/turbomind/engine/engine_config.h +++ b/src/turbomind/engine/engine_config.h @@ -29,6 +29,9 @@ struct EngineConfig { X(int, cache_prompt_boundary_skip, 1) \ X(std::string, cache_generation, "auto") \ X(bool, enable_metrics, false) \ + X(std::string, spec_method, "") \ + X(int, spec_num_draft_tokens, 0) \ + X(std::vector, spec_tap_layer_ids) \ X(int, num_tokens_per_iter, 0) \ X(int, max_prefill_iters, 1) \ X(int, async_, 0) \ diff --git a/src/turbomind/engine/gateway.cc b/src/turbomind/engine/gateway.cc index 160f582961..15f09e595a 100644 --- a/src/turbomind/engine/gateway.cc +++ b/src/turbomind/engine/gateway.cc @@ -32,7 +32,7 @@ void Gateway::push(std::shared_ptr r) { if (TM_UNLIKELY(!size_)) { TM_LOG_ERROR("No queues available for submitting the request"); - notify({[r = std::move(r)] { UpdateState(*r, Request::kNoQueue, 0); }}); + notify({MakeRequestSignal(std::move(r), Request::kNoQueue, 0)}); return; } const int rank = next_.fetch_add(1, std::memory_order_relaxed) % size_; @@ -80,9 +80,7 @@ void Gateway::cancel(std::shared_ptr r) { // {-1: canceled, 0: queued, 1: active} if (r->cancel_flag.exchange(-1, std::memory_order_acq_rel) == 0) { - notify({[r = std::move(r)] { // - UpdateState(*r, Request::kCancel, 0); - }}); + notify({MakeRequestSignal(std::move(r), Request::kCancel, 0)}); } else { // request is picked up by engine diff --git a/src/turbomind/engine/model.cc b/src/turbomind/engine/model.cc new file mode 100644 index 0000000000..1a1dc0db62 --- /dev/null +++ b/src/turbomind/engine/model.cc @@ -0,0 +1,60 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#include "src/turbomind/engine/model.h" + +#include + +#include "src/turbomind/models/model_weight.h" +#include "src/turbomind/models/speculative/speculative_model.h" + +namespace turbomind { + +Model::Model(std::unique_ptr model, + std::unique_ptr vision_model, + std::unique_ptr spec_model, + const EngineParam& param, + Context& context, + int phases): + target{std::move(model)}, + vision{std::move(vision_model)}, + spec{std::move(spec_model)}, + status{param.max_batch_size, phases}, + input_processor{param, + target->weights().hidden_units, + target->weights().data_type, + param.async_ ? 2 : 1, + spec != nullptr, + spec ? spec->policy().max_proposals() + 1 : 1, + spec ? spec->requires_successor_input_embeddings() : false}, + generation{kFloat32, + param.max_batch_size, + param.session_len, + target->weights().vocab_size, + target->weights().output->output_dim * context.comm.h_tp_group->n_ranks(), + target->weights().hidden_units, + target->weights().data_type, + context.comm.h_tp_group, + param.async_ ? 2 : 1, + spec ? &spec->policy() : nullptr, + param.enable_metrics}, + output_processor{*target, context.comm.h_tp_group->rank(), param.async_ ? 2 : 1} +{ +} + +Model::~Model() = default; + +void Model::Run(BatchOp op, int phase, TensorMap& env) +{ + std::apply( + [&](auto&&... modules) { + auto run = [&](auto* module) { + if (module) { + module->Run(op, phase, env); + } + }; + (run(modules), ...); + }, + OrderedComponents()); +} + +} // namespace turbomind diff --git a/src/turbomind/engine/model.h b/src/turbomind/engine/model.h new file mode 100644 index 0000000000..5f804935da --- /dev/null +++ b/src/turbomind/engine/model.h @@ -0,0 +1,68 @@ +// Copyright (c) OpenMMLab. All rights reserved. +#pragma once + +#include +#include + +#include "src/turbomind/engine/batch.h" +#include "src/turbomind/generation/generation.h" +#include "src/turbomind/models/batch_status.h" +#include "src/turbomind/models/input_processor.h" +#include "src/turbomind/models/language_model.h" +#include "src/turbomind/models/llama/context.h" +#include "src/turbomind/models/llama/llama_params.h" +#include "src/turbomind/models/output_processor.h" +#include "src/turbomind/models/vision_model.h" + +namespace turbomind { + +class SpeculativeModel; + +// The served model: the target, the optional vision and speculative models, +// and the input, output, generation, and status machinery around them — the +// whole that the engine thread's host operations and the executor's device +// steps drive. It owns its parts; the models are declared before the +// machinery so the machinery is destroyed first, and the speculative +// composition is resolved once from the speculator during construction. +// Constructed once at the +// engine, the composition root, and shared with the executor by reference. +// The generic batch-op fanout lives here: Run drives every module in the +// canonical order recorded once in OrderedComponents, skipping absent +// optional modules. Workflows that differ from the generic fanout — the +// prepare and forward steps — stay with the executor's device bracket. +class Model { +public: + Model(std::unique_ptr model, + std::unique_ptr vision_model, + std::unique_ptr spec_model, + const EngineParam& param, + Context& context, + int phases); + + // Defined in model.cc, where SpeculativeModel is complete, so this header + // never pulls in a speculative header. + ~Model(); + + // Generic batch-op fanout: drives every module in the canonical order. + void Run(BatchOp op, int phase, TensorMap& env); + + std::unique_ptr target; + std::unique_ptr vision; + std::unique_ptr spec; + + BatchStatus status; + InputProcessor input_processor; + Generation generation; + OutputProcessor output_processor; + +private: + // Canonical module order for the generic fanout; absent optional modules + // (vision, speculative) are skipped by the visitor. + auto OrderedComponents() + { + return std::make_tuple( + vision.get(), &status, &generation, &input_processor, target.get(), spec.get(), &output_processor); + } +}; + +} // namespace turbomind diff --git a/src/turbomind/engine/model_executor.cc b/src/turbomind/engine/model_executor.cc index 943a51b838..bd2fa9c1cf 100644 --- a/src/turbomind/engine/model_executor.cc +++ b/src/turbomind/engine/model_executor.cc @@ -1,143 +1,370 @@ +// Copyright (c) OpenMMLab. All rights reserved. #include "src/turbomind/engine/model_executor.h" #include +#include +#include +#include #include "src/turbomind/core/allocator.h" #include "src/turbomind/core/check.h" #include "src/turbomind/core/copy.h" +#include "src/turbomind/core/scope.h" #include "src/turbomind/engine/batch.h" +#include "src/turbomind/engine/model.h" +#include "src/turbomind/generation/target_verification.h" #include "src/turbomind/kernels/gemm/types.h" #include "src/turbomind/models/language_model.h" -#include "src/turbomind/models/llama/llama_utils.h" -#include "src/turbomind/models/vision_model.h" +#include "src/turbomind/models/llama/llama_kernels.h" +#include "src/turbomind/models/model_weight.h" +#include "src/turbomind/models/speculative/hidden_state_tap.h" +#include "src/turbomind/models/speculative/speculative_model.h" #include "src/turbomind/utils/anomaly_handler.h" #include "src/turbomind/utils/cuda_utils.h" - -// #include "dbg.h" +#include "src/turbomind/utils/nvtx_utils.h" namespace turbomind { -using std::shared_ptr; using std::unique_ptr; +struct DraftContext; + struct ModelExecutor::Impl { + struct TargetPass { + Tensor hidden_states; + Tensor pre_final_residual; + Tensor embedding_storage; + Tensor head_storage; + }; + + // The served model, owned by the engine's composition root and shared by + // reference. + Model& model_; - LanguageModel& model_; - VisionModel* vision_model_; // nullable: only set for VLM checkpoints - LlamaLinear& linear_; + LlamaLinear& linear_; const int device_id_; - Queue>& inbound_; - Queue>& outbound_; + Queue>& inbound_; + Queue>& outbound_; + + // Executor-internal scratch shared by the prepare step and the target pass. + Buffer_ autoreg_ids_; std::thread internal_thread_; - void InternalThreadEntry() - { - TM_FUNCTION_SCOPE(); - TM_CUDA_CHECK(cudaSetDevice(device_id_)); + Impl(Model& model, + const EngineParam& param, + Context& context, + int device_id, + Queue>& inbound, + Queue>& outbound); + + ~Impl(); + + // Ordinary and shared execution. + + void InternalThreadEntry(); + + static void RunCopies(std::vector& copies); + + void Run(BatchData& data); + + // Device-bracket steps, in execution order. + + // Prepare: non-forward component fanout plus the prepare-only cases — + // ordinary autoregressive-id exposure, the composed engine's explicit + // draft-input publish, and the executor-owned k-offsets buffer published + // before the target prepares and filled at forward in both compositions. + void PrepareStep(int phase, TensorMap& env); - Stream stream = Stream::create(); - Allocator h_alloc = Allocator(kCPU); - Allocator d_alloc = Allocator(stream, false); + // Forward: the visible composition branch (ADR 0002) — the speculative + // round when a speculator is composed, the ordinary target pass otherwise. + void ForwardStep(int phase, TensorMap& env); - AnomalyHandler::instance().Init(0, 1000, 0, 1000, stream.handle()); + // Unprep: non-forward component fanout. + void UnprepStep(int phase, TensorMap& env); - core::ContextGuard ctx{stream, h_alloc, d_alloc}; + // Shared target-pass step: publishes the symmetric buffer, then embeds, + // patches, and decodes the submitted rows. Publishing here holds the + // publish-before-decode ordering at one site for both forward routines. + // The tap arrives as an argument so the ordinary path holds no speculative + // branch (nullptr) while the speculative round supplies its own. + TargetPass RunTarget(int phase, TensorMap& env, HiddenStateTap* taps); - // Default GEMM workspace for everything dispatched on this stream; - // bound for the whole work loop, which is the outer-most scope that - // drives `linear_`. - gemm::Workspace workspace{stream.handle()}; - auto ws_lifetime = linear_.With(workspace); + // The ordinary forward: exactly one target pass, then logits and sampling. + void ForwardTargetPass(int phase, TensorMap& env); - unique_ptr d; + void RunVisionPass(int phase, TensorMap& env); - while (inbound_.pop(d)) { - TM_CHECK_NOTNULL(d); - core::Context::stream().Wait(d->ready); - Run(*d); - d->done.Record(core::Context::stream()); - outbound_.push(std::move(d)); - } + void Start(); - // Stream-ordered teardown: the frees run after the last batch's kernels. - workspace.Release(stream.handle()); + // Composed speculative execution. + + // The composed forward: one full speculative round owning the mixed batch. + void ForwardSpeculativeRound(int phase, TensorMap& env); + + DraftContext MakeDraftContext(const TargetPass& target, TensorMap& env); +}; + +ModelExecutor::Impl::Impl(Model& model, + const EngineParam& param, + Context& context, + int device_id, + Queue>& inbound, + Queue>& outbound): + model_{model}, + linear_{*context.linear}, + device_id_{device_id}, + inbound_{inbound}, + outbound_{outbound}, + autoreg_ids_{param.max_batch_size, kDEVICE} +{ +} + +ModelExecutor::Impl::~Impl() +{ + if (internal_thread_.joinable()) { + internal_thread_.join(); } +} - static void RunCopies(std::vector& copies) - { - for (const auto& c : copies) { - Copy(Buffer_{static_cast(c.src), static_cast(c.bytes), kDEVICE}, - Buffer_{static_cast(c.dst), static_cast(c.bytes), kDEVICE}); - } - copies.clear(); +void ModelExecutor::Impl::InternalThreadEntry() +{ + TM_FUNCTION_SCOPE(); + TM_CUDA_CHECK(cudaSetDevice(device_id_)); + + Stream stream = Stream::create(); + Allocator h_alloc = Allocator(kCPU); + Allocator d_alloc = Allocator(stream, false); + + AnomalyHandler::instance().Init(0, 1000, 0, 1000, stream.handle()); + + core::ContextGuard ctx{stream, h_alloc, d_alloc}; + + // Default GEMM workspace for everything dispatched on this stream; + // bound for the whole work loop, which is the outer-most scope that + // drives `linear_`. + gemm::Workspace workspace{stream.handle()}; + auto ws_lifetime = linear_.With(workspace); + + unique_ptr data; + while (inbound_.pop(data)) { + TM_CHECK_NOTNULL(data); + core::Context::stream().Wait(data->ready); + Run(*data); + data->done.Record(core::Context::stream()); + outbound_.push(std::move(data)); } - void Run(BatchData& d) - { - TM_FUNCTION_SCOPE(); + // Stream-ordered teardown: the frees run after the last batch's kernels. + workspace.Release(stream.handle()); +} + +void ModelExecutor::Impl::RunCopies(std::vector& copies) +{ + for (const auto& copy : copies) { + Copy(Buffer_{static_cast(copy.src), static_cast(copy.bytes), kDEVICE}, + Buffer_{static_cast(copy.dst), static_cast(copy.bytes), kDEVICE}); + } + copies.clear(); +} + +void ModelExecutor::Impl::Run(BatchData& data) +{ + TM_FUNCTION_SCOPE(); + + BatchCopy copy; + TensorMap env{{"batch", data.buf()}, {"copy", copy.buf()}}; - BatchCopy copy; - TensorMap env{{"batch", d.buf()}, {"copy", copy.buf()}}; + RunCopies(data.restore_copies); - // Restore copies first so kPrepare may post-process restored content - // (a module reset overrides whatever a whole-object restore wrote). - RunCopies(d.restore_copies); + PrepareStep(data.phase, env); + copy.Run(); - // Vision sub-graph runs before the language model in each phase so its - // env outputs (image embeddings, mrope tensors) are visible downstream. - if (vision_model_) { - vision_model_->Run(BatchOp::kPrepare, d.phase, env); - } - model_.Run(BatchOp::kPrepare, d.phase, env); - copy.Run(); + ForwardStep(data.phase, env); - if (vision_model_) { - vision_model_->Run(BatchOp::kForward, d.phase, env); - } - model_.Run(BatchOp::kForward, d.phase, env); + UnprepStep(data.phase, env); + copy.Run(); - model_.Run(BatchOp::kUnprep, d.phase, env); - copy.Run(); + RunCopies(data.publish_copies); - // Publication copies last: kUnprep is the module's final chance to - // finalize frontier contents before the snapshot. - RunCopies(d.publish_copies); + AnomalyHandler::instance().Summarize([](...) {}); + AnomalyHandler::instance().Reset(); +} - AnomalyHandler::instance().Summarize([](...) {}); - AnomalyHandler::instance().Reset(); +void ModelExecutor::Impl::PrepareStep(int phase, TensorMap& env) +{ + if (model_.vision) { + model_.vision->Run(BatchOp::kPrepare, phase, env); } - Impl(LanguageModel& model, - VisionModel* vision_model, - Context& context, - int device_id, - Queue>& inbound, - Queue>& outbound): - model_{model}, - vision_model_{vision_model}, - linear_{*context.linear}, - device_id_{device_id}, - inbound_{inbound}, - outbound_{outbound} - { + if (!model_.spec) { + env.emplace("autoreg_ids", autoreg_ids_); } - ~Impl() - { - if (internal_thread_.joinable()) { - internal_thread_.join(); - } + model_.status.Run(BatchOp::kPrepare, phase, env); + model_.generation.Run(BatchOp::kPrepare, phase, env); + if (model_.spec) { + model_.generation.Verification()->PublishDraftInputs(phase, env); + } + model_.input_processor.Run(BatchOp::kPrepare, phase, env); + + const int batch_size = env.at("batch").data()[0]->bsz; + env.produce("k_offsets", Buffer_{batch_size + 1, kDEVICE}); + + model_.target->Run(BatchOp::kPrepare, phase, env); + if (model_.spec) { + model_.spec->Run(BatchOp::kPrepare, phase, env); + } + model_.output_processor.Run(BatchOp::kPrepare, phase, env); +} + +void ModelExecutor::Impl::ForwardStep(int phase, TensorMap& env) +{ + if (model_.spec) { + return ForwardSpeculativeRound(phase, env); + } + ForwardTargetPass(phase, env); +} + +void ModelExecutor::Impl::UnprepStep(int phase, TensorMap& env) +{ + model_.Run(BatchOp::kUnprep, phase, env); +} + +void ModelExecutor::Impl::RunVisionPass(int phase, TensorMap& env) +{ + if (model_.vision) { + model_.vision->Run(BatchOp::kForward, phase, env); + } +} + +ModelExecutor::Impl::TargetPass ModelExecutor::Impl::RunTarget(int phase, TensorMap& env, HiddenStateTap* taps) +{ + const auto& batch = *env.at("batch").data()[0]; + if (batch.symm_buf) { + env.insert_or_assign("symm_buf", batch.symm_buf); + } + + auto& copy = *env.at("copy").data()[0]; + + Tensor residual = model_.target->Embed(env.at("input_ids").buffer(), {}, env); + TM_DEBUG_TENSOR(residual, "embeddings", 1); + + model_.input_processor.PatchEmbedding(phase, residual, copy, env); + copy.Run(); + + TargetPass out; + out.embedding_storage = residual; + + LanguageModel::DecoderInputs in; + in.residual = std::move(residual); + in.taps = taps; + in.selected_token_pos = env.consume("selected_token_pos").buffer(); + in.selected_hidden_buffer = env.try_consume("selected_normalized_hidden_buffer"); + + auto decoded = model_.target->RunDecoder(phase, in, env); + + out.hidden_states = std::move(decoded.selected_hidden); + out.pre_final_residual = std::move(decoded.pre_final_residual); + return out; +} + +void ModelExecutor::Impl::ForwardTargetPass(int phase, TensorMap& env) +{ + TM_FUNCTION_SCOPE(); + NvtxScope forward_scope{"targetExecutorForward"}; + + RunVisionPass(phase, env); + + const auto stream = core::Context::stream().handle(); + const auto& batch = *env.at("batch").data()[0]; + + PrefixSum(model_.status.SequenceLength().data(), batch.bsz, env.at("k_offsets").buffer().data(), stream); + + TargetPass target = RunTarget(phase, env, nullptr); + + model_.output_processor.OutputHiddenStatesAndLogits(phase, env, 2); + + target.head_storage = model_.target->Logits(target.hidden_states, {}, env); + env.produce("logits", target.head_storage); + + model_.output_processor.OutputHiddenStatesAndLogits(phase, env, 1); + + if (model_.status.GeneratingCount(phase)) { + model_.generation.Run(BatchOp::kForward, phase, env); + Copy(env.at("output_ids").buffer(), autoreg_ids_); + } +} + +DraftContext ModelExecutor::Impl::MakeDraftContext(const TargetPass& target, TensorMap& env) +{ + const auto& batch = *env.at("batch").data()[0]; + + DraftContext ctx; + ctx.batch_size = batch.bsz; + ctx.target_q_offsets = env.at("q_offsets").buffer(); + ctx.target_k_offsets = env.at("k_offsets").buffer(); + ctx.accept_len = env.at("accept_len").buffer(); + ctx.sequence_length = model_.status.SequenceLength(); + ctx.finished_on_entry = env.at("finished_on_entry").buffer(); + ctx.request_token_ids_ptrs = env.at("request_token_ids_ptrs").data(); + ctx.target_pre_final_residual = target.pre_final_residual; + ctx.embedding_storage = target.embedding_storage; + ctx.head_storage = target.head_storage; + ctx.input_processor = &model_.input_processor; + return ctx; +} + +void ModelExecutor::Impl::ForwardSpeculativeRound(int phase, TensorMap& env) +{ + TM_FUNCTION_SCOPE(); + NvtxScope forward_scope{"speculativeExecutorForward"}; + + RunVisionPass(phase, env); + + model_.input_processor.BuildTargetInputs(phase, env); + + const auto& batch = *env.at("batch").data()[0]; + cudaStream_t stream = core::Context::stream().handle(); + PrefixSum(env.at("target_key_lengths").buffer().data(), batch.bsz, + env.at("k_offsets").buffer().data(), stream); + + auto& verification = *model_.generation.Verification(); + env.produce("selected_normalized_hidden_buffer", + verification.SelectedHiddenBuffer(phase, env.at("selected_token_pos").size())); + + HiddenStateTap* tap = model_.spec->Tap(phase); + if (tap) { + tap = tap->Arm(phase, batch.local_token_num, stream); } - void Start() + TargetPass target = RunTarget(phase, env, tap); + + const int positions = model_.status.VerificationPositions(phase); + + verification.InitializeTargetVerification(phase, positions, env); + + target.head_storage = model_.target->Logits(target.hidden_states, {}, env); + verification.ProcessTargetBlock(phase, positions, target.head_storage, env); + verification.ClampSelectedSpan(phase, env); + + model_.target->CommitAcceptedState(phase, env.at("accept_len").buffer()); + { - internal_thread_ = std::thread(&Impl::InternalThreadEntry, this); + NvtxScope draft_scope{"SpeculativeModel::RunDraft"}; + model_.spec->RunDraft(phase, MakeDraftContext(target, env), env); } -}; + + verification.CommitAcceptedSpan(phase, model_.status.SequenceLength()); +} + +void ModelExecutor::Impl::Start() +{ + internal_thread_ = std::thread(&Impl::InternalThreadEntry, this); +} ModelExecutor::~ModelExecutor() = default; @@ -145,13 +372,13 @@ ModelExecutor::ModelExecutor() = default; ModelExecutor::ModelExecutor(ModelExecutor&&) noexcept = default; ModelExecutor& ModelExecutor::operator=(ModelExecutor&&) noexcept = default; -ModelExecutor::ModelExecutor(LanguageModel& model, - VisionModel* vision_model, - Context& context, - int device_id, - Queue>& inbound, - Queue>& outbound): - impl_{std::make_unique(model, vision_model, context, device_id, inbound, outbound)} +ModelExecutor::ModelExecutor(Model& model, + const EngineParam& param, + Context& context, + int device_id, + Queue>& inbound, + Queue>& outbound): + impl_{std::make_unique(model, param, context, device_id, inbound, outbound)} { } diff --git a/src/turbomind/engine/model_executor.h b/src/turbomind/engine/model_executor.h index 501267a488..9c7c22e678 100644 --- a/src/turbomind/engine/model_executor.h +++ b/src/turbomind/engine/model_executor.h @@ -1,4 +1,5 @@ // Copyright (c) OpenMMLab. All rights reserved. +#pragma once #include @@ -6,13 +7,14 @@ #include "src/turbomind/engine/batch.h" #include "src/turbomind/engine/queue.h" -#include "src/turbomind/models/language_model.h" -#include "src/turbomind/models/vision_model.h" #include "src/turbomind/models/llama/context.h" +#include "src/turbomind/models/llama/llama_params.h" namespace turbomind { +class Model; + // Model executor for auto-regressive language models, optionally // preceded by a per-batch ViT pass for VLM checkpoints. class ModelExecutor { @@ -28,8 +30,8 @@ class ModelExecutor { return static_cast(impl_); } - ModelExecutor(LanguageModel& model, - VisionModel* vision_model, // nullable + ModelExecutor(Model& model, + const EngineParam& param, Context& context, int device_id, Queue>& inbound, diff --git a/src/turbomind/engine/model_request.cc b/src/turbomind/engine/model_request.cc index f52c0d08d3..c7ece81735 100644 --- a/src/turbomind/engine/model_request.cc +++ b/src/turbomind/engine/model_request.cc @@ -16,12 +16,14 @@ namespace turbomind { -ModelRequest::ModelRequest(Gateway* gateway, DataType data_type, int session_len, int vocab_size, int hidden_dim): +ModelRequest::ModelRequest( + Gateway* gateway, DataType data_type, int session_len, int vocab_size, int hidden_dim, int speculative_tokens): gateway_{gateway}, data_type_{data_type}, session_len_{session_len}, + hidden_dim_{hidden_dim}, vocab_size_{vocab_size}, - hidden_dim_{hidden_dim} + speculative_tokens_{speculative_tokens} { } @@ -102,7 +104,7 @@ auto ModelRequest::Forward(InputParam param, std::function cb) -> Output auto state = std::make_shared(); - auto metrics = param.enable_metrics ? std::make_shared() : nullptr; + auto metrics = param.enable_metrics ? std::make_shared(speculative_tokens_) : nullptr; if (metrics) { metrics->enqueue_time.store(RequestMetrics::timestamp(), std::memory_order_relaxed); metrics->scheduled_time.store(0, std::memory_order_relaxed); diff --git a/src/turbomind/engine/model_request.h b/src/turbomind/engine/model_request.h index 7161c46d5e..de2fd93872 100644 --- a/src/turbomind/engine/model_request.h +++ b/src/turbomind/engine/model_request.h @@ -18,7 +18,8 @@ class ModelRequest { public: virtual ~ModelRequest() = default; - ModelRequest(Gateway* gateway, DataType data_type, int session_len, int vocab_size, int hidden_dim); + ModelRequest( + Gateway* gateway, DataType data_type, int session_len, int vocab_size, int hidden_dim, int speculative_tokens); // Cancel running request void Cancel(); @@ -52,6 +53,7 @@ class ModelRequest { const int session_len_; const int hidden_dim_; const int vocab_size_; + const int speculative_tokens_; uint64_t session_id_; diff --git a/src/turbomind/engine/request.cc b/src/turbomind/engine/request.cc index f7d416ce0f..bb93b57cec 100644 --- a/src/turbomind/engine/request.cc +++ b/src/turbomind/engine/request.cc @@ -44,21 +44,43 @@ std::ostream& operator<<(std::ostream& os, const GenerationConfig& c) return os; } -void UpdateState(Request& r, int status, int seq_len) +void UpdateState(Request& request, RequestState state) { try { - auto new_state = new RequestState{status, seq_len}; - auto old_state = r.state->exchange(new_state); - if (!old_state && r.forward_cb) { - r.forward_cb(); + auto next = new RequestState{std::move(state)}; + auto previous = request.state->exchange(next); + if (!previous && request.forward_cb) { + request.forward_cb(); } } catch (const std::exception& e) { - TM_LOG_ERROR("Error invoking callback for ({}): {}", r.id, e.what()); + TM_LOG_ERROR("Error invoking callback for ({}): {}", request.id, e.what()); } catch (...) { - TM_LOG_ERROR("Unknown error invoking callback for ({})", r.id); + TM_LOG_ERROR("Unknown error invoking callback for ({})", request.id); } } +std::function MakeRequestSignal(std::shared_ptr request, int status, int seq_len) +{ + RequestState state; + state.status = status; + state.seq_len = seq_len; + + if (request->metrics) { + auto& metrics = *request->metrics; + std::scoped_lock lock(metrics.spec_mutex); + + if (!metrics.num_accepted_tokens_per_pos.empty()) { + state.num_drafts = metrics.num_drafts; + state.num_draft_tokens = metrics.num_draft_tokens; + state.num_accepted_tokens = metrics.num_accepted_tokens; + state.num_accepted_tokens_per_pos = metrics.num_accepted_tokens_per_pos; + } + } + + return + [request = std::move(request), state = std::move(state)]() mutable { UpdateState(*request, std::move(state)); }; +} + } // namespace turbomind diff --git a/src/turbomind/engine/request.h b/src/turbomind/engine/request.h index 368b2cad4c..ea625ab7bd 100644 --- a/src/turbomind/engine/request.h +++ b/src/turbomind/engine/request.h @@ -9,6 +9,7 @@ #include #include #include +#include #include #include #include @@ -67,8 +68,14 @@ struct SessionParam { }; struct RequestState { - int status; - int seq_len; + int status{}; + int seq_len{}; + + int64_t num_drafts{}; + int64_t num_draft_tokens{}; + int64_t num_accepted_tokens{}; + + std::vector num_accepted_tokens_per_pos; }; struct AtomicRequestState { @@ -135,12 +142,49 @@ struct Request { std::shared_ptr matcher; }; -void UpdateState(Request& r, int status, int seq_len); +void UpdateState(Request& request, RequestState state); + +std::function MakeRequestSignal(std::shared_ptr request, int status, int seq_len); struct Sequence; struct MultiModalData; // defined in models/vision_model.h +struct SubmittedRow { + int input_len{}; + int history_len{}; + + int query_begin{}; + int query_count{}; + int key_capacity_end{}; + + int cache_write_begin{}; + int cache_write_end{}; + + bool generating{}; + bool autoregres{}; + + // Producer-set effects: assigned by the scheduler together with the row's + // geometry; consumers read values instead of classifying rows by mode (ADR 0003). + int verification_positions{}; // generating rows: query positions contributed to verification (K for speculative) + int min_grant{1}; // fewest query rows admission may grant partially + int inflight_input_delta{}; // completion effect on the sequence's inflight input length + int inflight_new_delta{}; // completion effect on the sequence's inflight new tokens + bool frontier_reanchor{}; // completed row re-anchors the resume frontier to the verified prefix + bool primes_proposals{}; // forward primes the first proposals (bootstrap fold) + + // Workload-shape selections over the effects (ADR 0003). + bool is_verification_row() const noexcept + { + return verification_positions > 1; + } + + bool is_extension_candidate() const noexcept + { + return is_verification_row() || primes_proposals; + } +}; + // The prefix-cache projection of one multimodal input: its token span and // content identity. The engine never sees MultiModalData / pixels. struct MultiModalSpan { @@ -193,11 +237,7 @@ struct Sequence { int seq_len = 0; // set at request init, updated per step - int input_len = 0; // set at schedule - int history_len = 0; // set at schedule from `resume_len` - - bool autoregres = false; // set at schedule, `seq_len` and `input_ids` taken from the engine - bool generating = false; // set at schedule + std::optional submitted; bool done = false; // set at cancel / update, is the request finished / canceled @@ -347,8 +387,9 @@ class Resource { public: virtual ~Resource() = default; - virtual int Test(const Sequence& s) const noexcept = 0; - virtual void Commit(const Sequence& s) noexcept = 0; + virtual int Test(const Sequence& s, const SubmittedRow& row) const noexcept = 0; + + virtual void Commit(const Sequence& s, const SubmittedRow& row) noexcept = 0; }; class ScheduleResources final: public Resource { @@ -362,11 +403,11 @@ class ScheduleResources final: public Resource { return ref; } - int Test(const Sequence& s) const noexcept override + int Test(const Sequence& s, const SubmittedRow& row) const noexcept override { int admitted = std::numeric_limits::max(); for (const auto& resource : resources_) { - const int next = resource->Test(s); + const int next = resource->Test(s, row); if (next == 0) { return 0; } @@ -375,10 +416,10 @@ class ScheduleResources final: public Resource { return admitted == std::numeric_limits::max() ? 0 : admitted; } - void Commit(const Sequence& s) noexcept override + void Commit(const Sequence& s, const SubmittedRow& row) noexcept override { for (const auto& resource : resources_) { - resource->Commit(s); + resource->Commit(s, row); } } @@ -390,18 +431,19 @@ class ForwardTokenResource final: public Resource { public: explicit ForwardTokenResource(int max_fwd_tokens) noexcept: max_fwd_tokens_{max_fwd_tokens} {} - int Test(const Sequence& s) const noexcept override + int Test(const Sequence&, const SubmittedRow& row) const noexcept override { - const int input_len = InputLen(s); - if (input_len <= 0 || max_fwd_tokens_ <= 0) { + const int q = row.query_count; + if (q <= 0 || max_fwd_tokens_ <= 0) { return 0; } - return std::min(input_len, max_fwd_tokens_); + + return max_fwd_tokens_ >= row.min_grant ? std::min(q, max_fwd_tokens_) : 0; } - void Commit(const Sequence& s) noexcept override + void Commit(const Sequence&, const SubmittedRow& row) noexcept override { - max_fwd_tokens_ -= s.input_len; + max_fwd_tokens_ -= row.query_count; } int remaining_tokens() const noexcept @@ -410,11 +452,6 @@ class ForwardTokenResource final: public Resource { } private: - static int InputLen(const Sequence& s) noexcept - { - return s.seq_len + s.inflight_new_tokens - s.inflight_input_len - s.resume_len; - } - int max_fwd_tokens_{}; }; diff --git a/src/turbomind/engine/scheduler.cc b/src/turbomind/engine/scheduler.cc index 847d9a0a3d..c92d2d338d 100644 --- a/src/turbomind/engine/scheduler.cc +++ b/src/turbomind/engine/scheduler.cc @@ -24,6 +24,17 @@ inline int InitialResumeUpperBound(const Sequence& s) return std::max(0, std::min(s.seq_len, context_len - 1)); } +// Ordinary-row effects, shared by row planning and the partial-admission clamp. +inline void SetOrdinaryEffects(SubmittedRow& row) +{ + row.verification_positions = row.generating; + row.min_grant = 1; + row.inflight_input_delta = row.input_len; + row.inflight_new_delta = row.generating; + row.frontier_reanchor = false; + row.primes_proposals = false; +} + // Clear per-pass planning buffers (alloc, restore, publish); involved_blocks persists. inline void ResetPassBuffers(Sequence& s) { @@ -213,7 +224,7 @@ struct GenStat { // within-pass facts. `bs` only where range math needs it. void LogAccept(const Sequence& s, int bs); void LogResume(const Sequence& s); -void LogDeferred(const Sequence& s, int bs, const Scheduler::ProducerConflict& c); +void LogDeferred(const Sequence& s, const SubmittedRow& candidate, int bs, const Scheduler::ProducerConflict& c); void LogPublished(const Sequence& s, int bs, const Scheduler::PublishStat& p); void LogFinalized(const Sequence& s, int bs, const GenStat& g); void LogCollision(const Sequence& s, CollisionSite site, int begin, int end); @@ -251,6 +262,7 @@ struct Scheduler::ScheduleState { Replay replay; // alloc/evict ops of the current phase size_t committed_replay_size{0}; // replay prefix from committed requests (phase 1) std::vector committed; + std::vector candidates; // one prepared forward per request std::vector pending_populate; // partial sibling node per request, nullptr = none std::vector pending_publish; // checkpoint publication intent per request bool has_optionals{false}; // any optional intent recorded => run phase 2 @@ -289,11 +301,15 @@ Scheduler::Scheduler(ObjectAllocator& alloc, const std::string& cache_prompt, int cache_prompt_boundary_skip, const std::string& cache_generation, + int session_len, + const SpeculativePolicy* policy, const int& is_warm_up): enable_prefix_caching_{enable_prefix_caching}, prompt_cache_mode_{ParseCacheMode(cache_prompt)}, cache_prompt_boundary_skip_{cache_prompt_boundary_skip < 1 ? 1 : cache_prompt_boundary_skip}, generation_cache_mode_{ParseCacheMode(cache_generation)}, + session_len_{session_len}, + policy_{policy}, is_warm_up_{is_warm_up}, alloc_{alloc}, registry_{std::move(registry)}, @@ -323,11 +339,10 @@ Scheduler::~Scheduler() } } -void Scheduler::EnsureBlocks(Sequence& s) +void Scheduler::EnsureBlocks(Sequence& s, int end) { const int bs = logical_.block_size(); - const int length = s.seq_len + s.inflight_new_tokens; - const int needed = (length + bs - 1) / bs; + const int needed = (end + bs - 1) / bs; while (static_cast(s.block_ids.size()) < needed) { const int i = static_cast(s.block_ids.size()); LogicalBlockPtr h = logical_.Create(i); @@ -517,7 +532,7 @@ void Scheduler::PlanResume(Sequence& s) s.resuming = true; - EnsureBlocks(s); + EnsureBlocks(s, s.seq_len + s.inflight_new_tokens); ResetPlanBuffers(s); @@ -667,7 +682,7 @@ void Scheduler::PlanContinue(Sequence& s) s.resuming = false; const int first_new = static_cast(s.block_ids.size()); - EnsureBlocks(s); + EnsureBlocks(s, s.seq_len + s.inflight_new_tokens); ResetPassBuffers(s); // per-pass buffers only; involved_blocks persists @@ -785,8 +800,7 @@ void Scheduler::Release(Sequence& s) s.resume_len = 0; s.filled_len = 0; s.readonly_block_num = 0; - s.input_len = 0; - s.history_len = 0; + s.submitted.reset(); } void Scheduler::Finalize(Sequence& s) @@ -1098,8 +1112,82 @@ void Scheduler::PlanRequests(ScheduleState& pass) } pass.committed.assign(pass.requests.size(), false); + pass.candidates.assign(pass.requests.size(), SubmittedRow{}); pass.pending_populate.assign(pass.requests.size(), nullptr); pass.pending_publish.assign(pass.requests.size(), PublishPlan{}); + + for (int i = 0; i < static_cast(pass.requests.size()); ++i) { + Sequence& s = *pass.requests[i]; + + const std::optional prior = s.submitted; + const std::optional outstanding = s.inflight != 0 ? prior : std::nullopt; + const bool was_generating = prior && prior->generating; + + const bool outstanding_speculative = + outstanding && outstanding->is_verification_row(); + const bool outstanding_final_prompt = + outstanding && !outstanding->is_verification_row() && outstanding->generating + && outstanding->key_capacity_end == s.prompt_len; + const bool no_outstanding_decode_row = + s.inflight == 0 && s.seq_len > s.prompt_len && s.filled_len == s.seq_len - 1 + && s.resume_len == s.filled_len; + + const int begin = s.resume_len + s.inflight_input_len; + const int context_end = s.seq_len + s.inflight_new_tokens; + + SubmittedRow& row = pass.candidates[i]; + + if (policy_ != nullptr + && (outstanding_speculative || outstanding_final_prompt || no_outstanding_decode_row)) { + row.history_len = s.resume_len; + row.generating = true; + row.autoregres = true; + + row.query_begin = outstanding_speculative + ? s.filled_len + outstanding->query_count + : (outstanding_final_prompt ? s.prompt_len : s.filled_len); + row.cache_write_begin = outstanding_speculative ? s.filled_len + 1 : row.query_begin; + + const RoundExtent extent = policy_->Extent({row.query_begin, s.prompt_len}); + row.query_count = extent.query_rows; + row.input_len = extent.query_rows; + row.key_capacity_end = row.query_begin + extent.query_rows; + row.cache_write_end = row.key_capacity_end + extent.private_tail; + + row.verification_positions = extent.query_rows; + row.min_grant = extent.query_rows; + row.inflight_input_delta = 0; + row.inflight_new_delta = 0; + row.frontier_reanchor = true; + row.primes_proposals = false; + } + else { + const int end = ClampForwardEnd(s, begin, context_end, context_end); + if (end <= begin) { + continue; + } + + row.history_len = s.resume_len; + row.input_len = end - begin; + row.query_begin = begin; + row.query_count = row.input_len; + row.key_capacity_end = end; + row.cache_write_begin = begin; + row.cache_write_end = end; + row.generating = end == context_end; + row.autoregres = s.is_active && was_generating && row.generating && row.query_count == 1; + + SetOrdinaryEffects(row); + + if (policy_ != nullptr && row.generating && row.key_capacity_end == s.prompt_len) { + if (auto bootstrap = policy_->Bootstrap(s.prompt_len); + bootstrap && bootstrap->min_session_len <= session_len_) { + row.cache_write_end = bootstrap->cache_write_end; + row.primes_proposals = true; + } + } + } + } } void Scheduler::RunRequiredAdmission(ScheduleState& pass, Resource& resource) @@ -1131,77 +1219,126 @@ void Scheduler::RunRequiredAdmission(ScheduleState& pass, Resource& resource) break; // would run on memory evicted from a higher-priority request } - const int admitted = resource.Test(s); + if (is_warm_up_ && s.inflight_input_len != 0) { + s.is_active = false; + continue; + } + + EvictingIterator evicting{evict_pos, pass.cutoff[i]}; + uint64_t evict_ts = max_evict_ts; + + const bool was_active = s.is_active; + const bool was_generating = s.submitted && s.submitted->generating; + const int context_end = s.seq_len + s.inflight_new_tokens; + SubmittedRow candidate = pass.candidates[i]; + + if (candidate.query_count <= 0) { + continue; + } + + const int admitted = resource.Test(s, candidate); if (admitted == 0) { TM_LOG_INFO("hit resource limit at {}/{}", i, pass.requests.size()); break; } - s.history_len = s.resume_len; + if (admitted < candidate.query_count) { + const int end = ClampForwardEnd(s, + candidate.query_begin, + candidate.query_begin + admitted, + context_end); + candidate.query_count = end - candidate.query_begin; + candidate.input_len = candidate.query_count; + candidate.key_capacity_end = end; + candidate.cache_write_end = end; + candidate.generating = end == context_end; + candidate.autoregres = was_active && was_generating && candidate.generating + && candidate.query_count == 1; + + SetOrdinaryEffects(candidate); - const int begin = s.resume_len + s.inflight_input_len; - const int ctx_end = s.seq_len + s.inflight_new_tokens; // == prompt_len for a fresh prefill + if (candidate.query_count <= 0) { + continue; + } + } - const int end = ClampForwardEnd(s, begin, begin + admitted, ctx_end); - const int len = end - begin; - if (len <= 0) { - continue; // nothing admitted this pass; CommitResults leaves it inactive + const int existing_end = std::min( + candidate.cache_write_end, + static_cast(s.block_ids.size()) * bs); + if (candidate.cache_write_begin < existing_end) { + if (const ProducerConflict conflict = + CheckProducers(s, candidate.cache_write_begin, existing_end); + conflict.producer) { + LogDeferred(s, candidate, bs, conflict); + continue; + } } - s.input_len = len; - // The publish decision is finalized in SetupPartialSiblings - // (prompt_boundary_node); the clamp lands a pass exactly on B iff it - // fired (an end past B implies begin >= B), so end == B identifies the - // prompt-boundary pass. - const bool at_prompt_boundary = s.prompt_boundary_node && end == s.prompt_boundary_pos; + const int required_end = candidate.cache_write_end; + std::vector staged_blocks; + std::vector staged_prefix; - if (const ProducerConflict conflict = CheckProducers(s, begin, end); conflict.producer) { - LogDeferred(s, bs, conflict); - continue; // deferred; CommitResults leaves it inactive + const int first_staged = static_cast(s.block_ids.size()); + const int needed_blocks = (required_end + bs - 1) / bs; + for (int index = first_staged; index < needed_blocks; ++index) { + LogicalBlockPtr block = logical_.Create(index); + block->prefix = cache_.Create(registry_.prefix().object_id(), block.get()); + staged_prefix.push_back(block->prefix.get()); + staged_blocks.push_back(std::move(block)); } - EvictingIterator evicting{evict_pos, pass.cutoff[i]}; - AllocatingIterator allocating{s.alloc_blocks}; + std::vector allocation_targets = s.alloc_blocks; + allocation_targets.insert(allocation_targets.end(), staged_prefix.begin(), staged_prefix.end()); - uint64_t evict_ts = 0; std::vector planned_now; - - bool ok = true; - while (allocating) { - bool success = allocating.Allocate(scratch, pass.planned, planned_now, pass.replay); - while (!success && evicting) { - evict_ts = evicting.Evict(scratch, pass.replay); - success = allocating.Allocate(scratch, pass.planned, planned_now, pass.replay); - } - if (!success) { - ok = false; - break; + bool allocated = true; + { + AllocatingIterator allocating{allocation_targets}; + while (allocating) { + bool success = allocating.Allocate(scratch, pass.planned, planned_now, pass.replay); + while (!success && evicting) { + evict_ts = evicting.Evict(scratch, pass.replay); + success = allocating.Allocate(scratch, pass.planned, planned_now, pass.replay); + } + if (!success) { + allocated = false; + break; + } } } - if (!ok) { // out of memory: roll back this request's planning, stop the pass - for (CacheBlock* b : planned_now) { - pass.planned.erase(b); + if (!allocated) { + for (CacheBlock* block : planned_now) { + pass.planned.erase(block); } TM_LOG_INFO("out of memory at {}/{}", i, pass.requests.size()); - break; // CommitResults leaves this and all later requests inactive + break; } - resource.Commit(s); + resource.Commit(s, candidate); + s.submitted = candidate; s.is_active = true; pass.committed[i] = true; pass.committed_replay_size = pass.replay.size(); - max_evict_ts = std::max(max_evict_ts, evict_ts); + max_evict_ts = evict_ts; evict_pos = evicting; - // Optional optimizations (allocated later, from inactive memory). One - // checkpoint per forward, routed by its end. - PlanPublication(pass, i, s, end, at_prompt_boundary); + for (LogicalBlockPtr& block : staged_blocks) { + CacheBlock* prefix = block->prefix.get(); + s.involved_blocks.push_back(prefix); + s.alloc_blocks.push_back(prefix); + s.block_ids.push_back(std::move(block)); + } + + const int target_end = candidate.query_begin + candidate.query_count; + if (policy_ == nullptr || target_end <= s.prompt_len) { + const bool at_prompt_boundary = s.prompt_boundary_node && target_end == s.prompt_boundary_pos; + PlanPublication(pass, i, s, target_end, at_prompt_boundary); + } - SetProducers(s, begin, end); + SetProducers(s, candidate.cache_write_begin, candidate.cache_write_end); if (s.resuming) { - // emit here so a producer's resume precedes any later consumer's defer log LogResume(s); } } @@ -1324,13 +1461,14 @@ void Scheduler::CommitResults(ScheduleState& pass) // branch; do not zero these fields at reject sites. if (!pass.committed[i]) { s.is_active = false; - s.input_len = 0; - s.history_len = 0; s.publish_target = nullptr; s.publish_end = 0; s.alloc_blocks.clear(); s.restore_copies.clear(); s.publish_copies.clear(); + if (s.inflight == 0) { + s.submitted.reset(); + } continue; } @@ -1345,8 +1483,9 @@ void Scheduler::CommitResults(ScheduleState& pass) s.filled_len = s.resume_len; } - const int begin = s.history_len + s.inflight_input_len; - const int end = begin + s.input_len; + const SubmittedRow& row = *s.submitted; + const int begin = row.query_begin; + const int end = begin + row.query_count; bool ckpt_published = false; @@ -1382,7 +1521,7 @@ void Scheduler::CommitResults(ScheduleState& pass) // Content is guaranteed to be produced by this iteration (device // execution is in submission order); no point deferring to Update(). - PublishStat pub = MarkProduced(s, begin, end); + PublishStat pub = MarkProduced(s, row.cache_write_begin, row.cache_write_end); pub.forked = pass.pending_populate[i] != nullptr; pub.ckpt = ckpt_published; LogPublished(s, bs, pub); @@ -1469,30 +1608,31 @@ void LogAccept(const Sequence& s, int bs) void LogResume(const Sequence& s) { auto msg = [&] { - const int begin = s.history_len + s.inflight_input_len; - const int end = begin + s.input_len; - const int total = s.seq_len + s.inflight_new_tokens; - const int pct = total > 0 ? 100 * s.history_len / total : 0; + const SubmittedRow& row = *s.submitted; + const int begin = row.query_begin; + const int end = begin + row.query_count; + const int total = s.seq_len + s.inflight_new_tokens; + const int pct = total > 0 ? 100 * row.history_len / total : 0; return fmt::format("req {} (uid {}) resume [0,{}) {} blk ro ({}%) source={} | computed [{},{}) {} tok", s.req->id, s.req->unique_id, - s.history_len, + row.history_len, s.readonly_block_num, pct, ResumeSourceName(s.resume_source), begin, end, - s.input_len); + row.input_len); }; TM_LOG(kCacheLogLevel, msg()); } -void LogDeferred(const Sequence& s, int bs, const Scheduler::ProducerConflict& c) +void LogDeferred(const Sequence& s, const SubmittedRow& candidate, int bs, const Scheduler::ProducerConflict& c) { auto msg = [&] { - const int begin = s.resume_len + s.inflight_input_len; - const int end = begin + s.input_len; + const int begin = candidate.query_begin; + const int end = begin + candidate.query_count; const int b0 = std::max(begin, c.block * bs); const int b1 = std::min(end, (c.block + 1) * bs); return fmt::format("req {} (uid {}) deferred: tok [{},{}) held by producer uid {}", @@ -1511,9 +1651,10 @@ void LogPublished(const Sequence& s, int bs, const Scheduler::PublishStat& p) return; } auto msg = [&] { - const int end = s.history_len + s.inflight_input_len + s.input_len; // forward end this pass - std::string body; - auto add = [&](std::string c) { body += body.empty() ? c : ", " + c; }; + const SubmittedRow& row = *s.submitted; + const int end = row.query_begin + row.query_count; + std::string body; + auto add = [&](std::string c) { body += body.empty() ? c : ", " + c; }; if (p.reusable_blocks > 0) { add(fmt::format("prefix [{},{}) ({} blk)", p.start, p.end, p.reusable_blocks)); } diff --git a/src/turbomind/engine/scheduler.h b/src/turbomind/engine/scheduler.h index bf7eaa1f7e..07c11aeb2c 100644 --- a/src/turbomind/engine/scheduler.h +++ b/src/turbomind/engine/scheduler.h @@ -12,6 +12,7 @@ #include "src/turbomind/engine/prefix_trie.h" #include "src/turbomind/engine/request.h" #include "src/turbomind/memory/object.h" +#include "src/turbomind/models/speculative/speculative_policy.h" #define TM_SCHED_PROFILE 0 @@ -93,6 +94,8 @@ class Scheduler { const std::string& cache_prompt, int cache_prompt_boundary_skip, const std::string& cache_generation, + int session_len, + const SpeculativePolicy* policy, const int& is_warm_up); ~Scheduler(); @@ -204,7 +207,7 @@ class Scheduler { // dedup intent across requests. void PlanPublication(ScheduleState& pass, int i, Sequence& s, int end, bool at_prompt_boundary); - void EnsureBlocks(Sequence& s); + void EnsureBlocks(Sequence& s, int end); bool PrefixEligible(const Sequence& s) const noexcept; bool CheckpointPublicationEligible() const noexcept; @@ -216,6 +219,8 @@ class Scheduler { CacheMode prompt_cache_mode_{CacheMode::kAuto}; int cache_prompt_boundary_skip_{1}; CacheMode generation_cache_mode_{CacheMode::kAuto}; + const int session_len_; + const SpeculativePolicy* policy_; const int& is_warm_up_; ObjectAllocator& alloc_; // owned by Engine; also used outside the scheduler CacheRegistry registry_; // owned: registration is closed before construction diff --git a/src/turbomind/generation/CMakeLists.txt b/src/turbomind/generation/CMakeLists.txt index 4c1c2313f0..6b58412083 100644 --- a/src/turbomind/generation/CMakeLists.txt +++ b/src/turbomind/generation/CMakeLists.txt @@ -22,6 +22,7 @@ set_property(TARGET guided_decoding PROPERTY POSITION_INDEPENDENT_CODE ON) add_library(generation STATIC generation.cc + target_verification.cc logits_processor.cc sampling.cc stop_criteria.cc) @@ -37,3 +38,6 @@ target_link_libraries(generation PUBLIC guided_decoding memory_utils CUDA::cudart) +target_link_libraries(generation PRIVATE + speculative_sampling_kernels + speculative_sequence_kernels) diff --git a/src/turbomind/generation/generation.cc b/src/turbomind/generation/generation.cc index 9a6588d777..891c01ee99 100644 --- a/src/turbomind/generation/generation.cc +++ b/src/turbomind/generation/generation.cc @@ -1,3 +1,4 @@ +// Copyright (c) OpenMMLab. All rights reserved. #include @@ -10,6 +11,7 @@ #include "src/turbomind/engine/batch.h" #include "src/turbomind/engine/request.h" +#include "src/turbomind/generation/generation_impl.h" #include "src/turbomind/generation/guided_decoding.h" #include "src/turbomind/generation/logits_processor.h" #include "src/turbomind/generation/sampling.h" @@ -18,9 +20,6 @@ #include "src/turbomind/kernels/sampling_topk_kernels.h" // InitializeRandomStates #include "src/turbomind/models/llama/llama_kernels.h" // invokePadLastTokenIds -#include "src/turbomind/utils/cuda_utils.h" - -// #include "dbg.h" namespace turbomind { @@ -28,125 +27,115 @@ using std::unique_ptr; using std::shared_ptr; using std::vector; -struct GenerationData { - Buffer_ random_seed; - Buffer_ random_init; - Buffer_ random_state_indices; - Buffer_ max_seq_len; - Buffer_ token_ids_ptrs; - Buffer_ output_ids; - - bool random_init_needed; - int generation_size; -}; - -struct Generation::Impl { - - // child modules - unique_ptr logits_processor_; - unique_ptr sampling_; - shared_ptr stop_criteria_; - unique_ptr guided_decoding_; - - // persistent - Tensor_ token_ids_; - Tensor_ random_states_; - - // scheduling states - vector free_token_rows_; - vector free_random_state_rows_; - - // immutable states - Buffer_ output_ids_; - - std::vector> data_; - - // staging buffers - Buffer_ random_seed_buf_; - Buffer_ random_init_buf_; - Buffer_ random_state_indices_buf_; - Buffer_ token_ids_ptrs_buf_; - Buffer_ token_ids_buf_; - Buffer_ output_ids_buf_; - - const int max_batch_size_; - const int session_len_; - - int* RowPtr(int row) - { - TM_CHECK_GE(row, 0); - TM_CHECK_LT(row, max_batch_size_); - return token_ids_.data() + row * token_ids_.stride(0); +Generation::Impl::Impl(DataType dtype, + int max_batch_size, + int session_len, + int vocab_size, + int vocab_size_padded, + int hidden_units, + DataType hidden_dtype, + const comm::HostComm& tp_group, + int phases, + const SpeculativePolicy* policy, + bool enable_metrics): + max_batch_size_{max_batch_size}, + session_len_{session_len}, + token_row_width_{session_len + (policy ? policy->token_row_tail() : 0)}, + policy_{policy}, + shared_{logits_processor_, + sampling_, + stop_criteria_, + random_states_, + token_ids_, + max_batch_size_, + token_row_width_, + token_ids_ptrs_buf_, + data_} +{ + TM_CHECK_EQ(dtype, kFloat32); + BaseGenerationParam base{max_batch_size, vocab_size, vocab_size_padded}; + const int verification_capacity = policy ? policy->max_proposals() + 1 : 1; + const int parameter_capacity = max_batch_size_ * verification_capacity; + logits_processor_ = + std::make_unique(base, phases, policy != nullptr, parameter_capacity); + sampling_ = std::make_unique(base, phases, tp_group->rank(), parameter_capacity); + stop_criteria_ = std::make_unique(base, phases); + guided_decoding_ = std::make_unique(base, tp_group, phases); + + static_assert(sizeof(curandState_t) % alignof(curandState_t) == 0); + random_states_ = {{max_batch_size_, (int)sizeof(curandState_t)}, kDEVICE}; + token_ids_ = {{max_batch_size_, token_row_width_}, kDEVICE}; + output_ids_ = {max_batch_size_, kDEVICE}; + for (int i = 0; i < max_batch_size_; ++i) { + free_token_rows_.push_back(i); + free_random_state_rows_.push_back(i); } - Impl(DataType dtype, - int max_batch_size, - int session_len, - int vocab_size, - int vocab_size_padded, - const comm::HostComm& tp_group, - int phases): - max_batch_size_{max_batch_size}, session_len_{session_len} - { - TM_CHECK_EQ(dtype, kFloat32); - BaseGenerationParam base{max_batch_size, vocab_size, vocab_size_padded}; - logits_processor_ = std::make_unique(base, phases); - sampling_ = std::make_unique(base, phases, tp_group->rank()); - stop_criteria_ = std::make_unique(base, phases); - guided_decoding_ = std::make_unique(base, tp_group, phases); - - static_assert(sizeof(curandState_t) % alignof(curandState_t) == 0); - random_states_ = {{max_batch_size_, (int)sizeof(curandState_t)}, kDEVICE}; - token_ids_ = {{max_batch_size_, session_len_}, kDEVICE}; - output_ids_ = {max_batch_size_, kDEVICE}; - for (int i = 0; i < max_batch_size_; ++i) { - free_token_rows_.push_back(i); - free_random_state_rows_.push_back(i); - } - - random_seed_buf_ = {max_batch_size_, kCPUpinned}; - random_init_buf_ = {max_batch_size_, kCPUpinned}; - random_state_indices_buf_ = {max_batch_size_, kCPUpinned}; + random_seed_buf_ = {max_batch_size_, kCPUpinned}; + random_init_buf_ = {max_batch_size_, kCPUpinned}; + random_state_indices_buf_ = {max_batch_size_, kCPUpinned}; - token_ids_ptrs_buf_ = {max_batch_size_, kCPUpinned}; - token_ids_buf_ = {max_batch_size_ * (ssize_t)session_len_, kCPUpinned}; + token_ids_ptrs_buf_ = {parameter_capacity, kCPUpinned}; + token_ids_buf_ = {max_batch_size_ * (ssize_t)session_len_, kCPUpinned}; + output_ids_buf_ = {max_batch_size_, kCPUpinned}; + request_to_generation_row_offsets_buf_ = {max_batch_size_ + 1, kCPUpinned}; - output_ids_buf_ = {max_batch_size_, kCPUpinned}; + for (int i = 0; i < phases; ++i) { + auto d = std::make_unique(); - for (int i = 0; i < phases; ++i) { - auto d = std::make_unique(); + d->random_seed = empty_like(random_seed_buf_, kDEVICE); + d->random_init = empty_like(random_init_buf_, kDEVICE); + d->random_state_indices = empty_like(random_state_indices_buf_, kDEVICE); + d->token_ids_ptrs = empty_like(token_ids_ptrs_buf_, kDEVICE); + d->request_to_generation_row_offsets = {max_batch_size_ + 1, kDEVICE}; + d->output_ids = empty_like(output_ids_, kDEVICE); - d->random_seed = empty_like(random_seed_buf_, kDEVICE); - d->random_init = empty_like(random_init_buf_, kDEVICE); - d->random_state_indices = empty_like(random_state_indices_buf_, kDEVICE); - d->token_ids_ptrs = empty_like(token_ids_ptrs_buf_, kDEVICE); - d->output_ids = empty_like(output_ids_, kDEVICE); + data_.push_back(std::move(d)); + } - data_.push_back(std::move(d)); - } + if (policy_) { + verification_ = std::make_unique( + shared_, *policy_, hidden_units, hidden_dtype, enable_metrics, phases); } +} - void Setup(int phase, TensorMap& env) - { - TM_FUNCTION_SCOPE(); - auto& d = *data_.at(phase); +void Generation::Impl::Setup(int phase, TensorMap& env) +{ + TM_FUNCTION_SCOPE(); + auto& d = *data_.at(phase); - auto& copy = *env.at("copy").data()[0]; + auto& copy = *env.at("copy").data()[0]; - Buffer_ rc = env.at("requests").buffer(); + Buffer_ rc = env.at("requests").buffer(); - // random states - d.random_init_needed = false; - std::fill_n(random_init_buf_.data(), max_batch_size_, false); + // random states + d.random_init_needed = false; + std::fill_n(random_init_buf_.data(), max_batch_size_, false); - int* token_ids_buf = token_ids_buf_.data(); - int generation_size = 0; - for (int i = 0; i < rc.size(); ++i) { - auto& c = *rc[i]; - if (!c.generating) { - continue; - } + int* token_ids_buf = token_ids_buf_.data(); + int generation_size = 0; + request_to_generation_row_offsets_buf_[0] = 0; + for (int i = 0; i < rc.size(); ++i) { + auto& c = *rc[i]; + const SubmittedRow& submitted = *c.submitted; + // An eagerly allocated row also serves a prompt whose extent the + // policy extends (the bootstrapping forward writes proposals into it). + const bool needs_row = submitted.generating || (policy_ && policy_->needs_prompt_token_row(c.prompt_len)); + + if (needs_row && c.generation_token_ids_row < 0) { + TM_CHECK(!free_token_rows_.empty()); + + c.generation_token_ids_row = free_token_rows_.back(); + free_token_rows_.pop_back(); + + auto* dst = shared_.RowPtr(c.generation_token_ids_row); + std::copy_n(c.token_ids, c.seq_len, token_ids_buf); + copy(token_ids_buf, c.seq_len, dst); + token_ids_buf += c.seq_len; + } + + if (submitted.generating) { if (c.generation_random_state_row < 0) { TM_CHECK(!free_random_state_rows_.empty()); @@ -158,145 +147,142 @@ struct Generation::Impl { d.random_init_needed = true; } - if (c.generation_token_ids_row < 0) { - TM_CHECK(!free_token_rows_.empty()); - - c.generation_token_ids_row = free_token_rows_.back(); - free_token_rows_.pop_back(); - - auto* dst = RowPtr(c.generation_token_ids_row); - std::copy_n(c.token_ids, c.seq_len, token_ids_buf); - copy(token_ids_buf, c.seq_len, dst); - token_ids_buf += c.seq_len; - } - random_state_indices_buf_[generation_size] = c.generation_random_state_row; - token_ids_ptrs_buf_[generation_size++] = RowPtr(c.generation_token_ids_row); + token_ids_ptrs_buf_[generation_size] = shared_.RowPtr(c.generation_token_ids_row); + ++generation_size; } - if (d.random_init_needed) { - copy(random_init_buf_, max_batch_size_, d.random_init); - copy(random_seed_buf_, max_batch_size_, d.random_seed); - } + request_to_generation_row_offsets_buf_[i + 1] = generation_size; + } + if (d.random_init_needed) { + copy(random_init_buf_, max_batch_size_, d.random_init); + copy(random_seed_buf_, max_batch_size_, d.random_seed); + } + if (!verification_) { copy(token_ids_ptrs_buf_, generation_size, d.token_ids_ptrs); - copy(random_state_indices_buf_, generation_size, d.random_state_indices); - d.generation_size = generation_size; - // dbg(d.generation_size); - - logits_processor_->Setup(phase, env); - sampling_->Setup(phase, env); - stop_criteria_->Setup(phase, env); - guided_decoding_->Setup(phase, env); } + else { + copy(request_to_generation_row_offsets_buf_, rc.size() + 1, d.request_to_generation_row_offsets); + } + copy(random_state_indices_buf_, generation_size, d.random_state_indices); + d.request_count = rc.size(); + d.generation_size = generation_size; - void Del(TensorMap& env) - { - Buffer_ rc = env.at("requests").buffer(); + logits_processor_->Setup(phase, env); + sampling_->Setup(phase, env); + stop_criteria_->Setup(phase, env); + guided_decoding_->Setup(phase, env); - for (int i = 0; i < rc.size(); ++i) { - auto& token_row = rc[i]->generation_token_ids_row; - if (token_row >= 0) { - free_token_rows_.push_back(token_row); - token_row = -1; - } + if (verification_) { + verification_->Setup(phase, env); + } +} - auto& random_row = rc[i]->generation_random_state_row; - if (random_row >= 0) { - free_random_state_rows_.push_back(random_row); - random_row = -1; - } +void Generation::Impl::Del(TensorMap& env) +{ + Buffer_ rc = env.at("requests").buffer(); + + for (int i = 0; i < rc.size(); ++i) { + auto& token_row = rc[i]->generation_token_ids_row; + if (token_row >= 0) { + free_token_rows_.push_back(token_row); + token_row = -1; } - } - void Prepare(int phase, TensorMap& env) - { - TM_FUNCTION_SCOPE(); - (void)phase; - (void)env; + auto& random_row = rc[i]->generation_random_state_row; + if (random_row >= 0) { + free_random_state_rows_.push_back(random_row); + random_row = -1; + } } +} - void Unprep(int phase, TensorMap& env) - { - TM_FUNCTION_SCOPE(); - auto& d = *data_.at(phase); - auto& b = *env.at("batch").data()[0]; - auto& copy = *env.at("copy").data()[0]; +void Generation::Impl::Unprep(int phase, TensorMap& env) +{ + TM_FUNCTION_SCOPE(); + auto& d = *data_.at(phase); + auto& b = *env.at("batch").data()[0]; + auto& copy = *env.at("copy").data()[0]; + if (!verification_) { copy(output_ids_, b.bsz, d.output_ids); } +} - void Fetch(int phase, TensorMap& env) - { - TM_FUNCTION_SCOPE(); - auto& d = *data_.at(phase); - auto& copy = *env.at("copy").data()[0]; +void Generation::Impl::Fetch(int phase, TensorMap& env) +{ + TM_FUNCTION_SCOPE(); + auto& d = *data_.at(phase); + auto& copy = *env.at("copy").data()[0]; + if (verification_) { + verification_->Fetch(phase, env); + } + else { copy(d.output_ids, d.output_ids.size(), output_ids_buf_); env.produce("output_ids", output_ids_buf_); sampling_->Fetch(phase, env); } +} - void Update(int phase, TensorMap& env) - { - TM_FUNCTION_SCOPE(); - sampling_->Update(phase, env); - } - - void Forward(int phase, TensorMap& env) - { - TM_FUNCTION_SCOPE(); - auto& d = *data_.at(phase); +void Generation::Impl::Update(int phase, TensorMap& env) +{ + TM_FUNCTION_SCOPE(); + sampling_->Update(phase, env); +} - const auto stream = core::Context::stream().handle(); +void Generation::Impl::Forward(int phase, TensorMap& env) +{ + TM_FUNCTION_SCOPE(); + auto& d = *data_.at(phase); - if (d.random_init_needed) { - InitializeRandomStates((curandState_t*)random_states_.raw_data(), - d.random_seed.data(), - d.random_init.data(), - max_batch_size_, - stream); - } + const auto stream = core::Context::stream().handle(); - env.emplace("output_ids", output_ids_); // out - env.emplace("curand_state", random_states_); // inout + if (d.random_init_needed) { + InitializeRandomStates((curandState_t*)random_states_.raw_data(), + d.random_seed.data(), + d.random_init.data(), + max_batch_size_, + stream); + } - if (const int gs = d.generation_size) { + env.emplace("output_ids", output_ids_); // out + env.emplace("curand_state", random_states_); // inout - env.emplace("token_ids_ptrs", d.token_ids_ptrs.slice(0, gs)); - env.emplace("curand_state_indices", d.random_state_indices.slice(0, gs)); + if (const int gs = d.generation_size) { - auto logits = env.consume("logits"); + env.emplace("token_ids_ptrs", d.token_ids_ptrs.slice(0, gs)); + env.emplace("curand_state_indices", d.random_state_indices.slice(0, gs)); - if (logits.dtype() != kFloat32) { - auto tmp = empty_like(logits, kFloat32); - TM_SCOPE_CALL(invokeCastFloat2D(logits, tmp, stream)); - logits = std::move(tmp); - } + auto logits = env.consume("logits"); - env.produce("logits", logits.slice(0, gs)); + if (logits.dtype() != kFloat32) { + auto tmp = empty_like(logits, kFloat32); + TM_SCOPE_CALL(invokeCastFloat2D(logits, tmp, stream)); + logits = std::move(tmp); + } - Buffer_ output_pos{max_batch_size_, kDEVICE}; - Copy(env.at("sequence_length").buffer(), gs, output_pos); + env.produce("logits", logits.slice(0, gs)); - logits_processor_->Forward(phase, env); + logits_processor_->Forward(phase, env); - guided_decoding_->FillMask(phase, env); - guided_decoding_->ApplyMask(phase, env); + guided_decoding_->FillMask(phase, env); + guided_decoding_->ApplyMask(phase, env); - sampling_->Forward(phase, env); + sampling_->Forward(phase, env); - guided_decoding_->ScheduleUpdate(phase, env); + guided_decoding_->ScheduleUpdate(phase, env); - AppendTokenIds(d.token_ids_ptrs.data(), output_ids_.data(), output_pos.data(), gs, stream); + invokeAppendOneTokenAndAdvanceSequence( + d.token_ids_ptrs.data(), output_ids_.data(), env.at("sequence_length").data(), gs, stream); - stop_criteria_->Forward(phase, env); + stop_criteria_->Forward(phase, env); - guided_decoding_->FinishUpdate(phase, env); - } + guided_decoding_->FinishUpdate(phase, env); } -}; +} Generation::~Generation() = default; @@ -305,9 +291,23 @@ Generation::Generation(DataType dtype, int session_len, int vocab_size, int vocab_size_padded, + int hidden_units, + DataType hidden_dtype, const comm::HostComm& tp_group, - int phases): - impl_{std::make_unique(dtype, max_batch_size, session_len, vocab_size, vocab_size_padded, tp_group, phases)} + int phases, + const SpeculativePolicy* policy, + bool enable_metrics): + impl_{std::make_unique(dtype, + max_batch_size, + session_len, + vocab_size, + vocab_size_padded, + hidden_units, + hidden_dtype, + tp_group, + phases, + policy, + enable_metrics)} { } @@ -319,9 +319,6 @@ void Generation::Run(BatchOp op, int phase, TensorMap& env) else if (op == BatchOp::kDel) { return impl_->Del(env); } - else if (op == BatchOp::kPrepare) { - return impl_->Prepare(phase, env); - } else if (op == BatchOp::kForward) { return impl_->Forward(phase, env); } @@ -336,4 +333,9 @@ void Generation::Run(BatchOp op, int phase, TensorMap& env) } } +TargetVerification* Generation::Verification() noexcept +{ + return impl_->verification_.get(); +} + } // namespace turbomind diff --git a/src/turbomind/generation/generation.h b/src/turbomind/generation/generation.h index 7261d9dc70..f25e418e5f 100644 --- a/src/turbomind/generation/generation.h +++ b/src/turbomind/generation/generation.h @@ -1,11 +1,11 @@ - - +// Copyright (c) OpenMMLab. All rights reserved. #pragma once #include #include "src/turbomind/core/core.h" #include "src/turbomind/engine/batch.h" +#include "src/turbomind/models/speculative/speculative_policy.h" namespace turbomind { @@ -13,22 +13,30 @@ namespace comm { class HostComm; } -struct GenerationData; +class TargetVerification; class Generation { public: ~Generation(); - Generation(DataType data_type, // + Generation(DataType data_type, int max_batch_size, int session_len, int vocab_size, int vocab_size_padded, + int hidden_units, + DataType hidden_dtype, const comm::HostComm& tp_group, - int phases); + int phases, + const SpeculativePolicy* policy, + bool enable_metrics); void Run(BatchOp op, int phase, TensorMap& env); + // The composed engine's verification component, holding the speculative + // round's verification steps and buffers. Null in target-only engines. + TargetVerification* Verification() noexcept; + private: struct Impl; diff --git a/src/turbomind/generation/generation_impl.h b/src/turbomind/generation/generation_impl.h new file mode 100644 index 0000000000..5987a56294 --- /dev/null +++ b/src/turbomind/generation/generation_impl.h @@ -0,0 +1,124 @@ +// Copyright (c) OpenMMLab. All rights reserved. +#pragma once + +#include +#include + +#include "src/turbomind/core/check.h" +#include "src/turbomind/core/core.h" +#include "src/turbomind/engine/batch.h" +#include "src/turbomind/generation/generation.h" +#include "src/turbomind/generation/target_verification.h" + +namespace turbomind { + +class LogitsProcessor; +class Sampling; +class StopCriteria; +class GuidedDecoding; + +// Per-phase shared generation state: rows, random states, and the request to +// generation-row topology used by both the ordinary sampling lifecycle and the +// verification component. +struct GenerationData { + Buffer_ random_seed; + Buffer_ random_init; + Buffer_ random_state_indices; + Buffer_ token_ids_ptrs; + Buffer_ request_to_generation_row_offsets; + Buffer_ output_ids; + + bool random_init_needed; + int request_count; + int generation_size; +}; + +// The generation-module state the verification component borrows. Reference +// bundle bound once at construction; the referenced members outlive it. +struct GenerationShared { + std::unique_ptr& logits_processor; + std::unique_ptr& sampling; + std::shared_ptr& stop_criteria; + + Tensor_& random_states; + Tensor_& token_ids; + const int& max_batch_size; + const int& token_row_width; + + Buffer_& token_ids_ptrs_buf; + + std::vector>& data; + + int* RowPtr(int row) const + { + TM_CHECK_GE(row, 0); + TM_CHECK_LT(row, max_batch_size); + return token_ids.data() + row * token_ids.stride(0); + } +}; + +struct Generation::Impl { + + // child modules + std::unique_ptr logits_processor_; + std::unique_ptr sampling_; + std::shared_ptr stop_criteria_; + std::unique_ptr guided_decoding_; + + // persistent + Tensor_ token_ids_; + Tensor_ random_states_; + + // scheduling states + std::vector free_token_rows_; + std::vector free_random_state_rows_; + + // immutable states + Buffer_ output_ids_; + + std::vector> data_; + + // staging buffers + Buffer_ random_seed_buf_; + Buffer_ random_init_buf_; + Buffer_ random_state_indices_buf_; + Buffer_ token_ids_ptrs_buf_; + Buffer_ token_ids_buf_; + Buffer_ output_ids_buf_; + Buffer_ request_to_generation_row_offsets_buf_; + + const int max_batch_size_; + const int session_len_; + const int token_row_width_; + const SpeculativePolicy* const policy_; + + GenerationShared shared_; + + // Present only when a speculative policy is (the composed engine's + // verification half); the ordinary lifecycle below never enters it. + std::unique_ptr verification_; + + Impl(DataType dtype, + int max_batch_size, + int session_len, + int vocab_size, + int vocab_size_padded, + int hidden_units, + DataType hidden_dtype, + const comm::HostComm& tp_group, + int phases, + const SpeculativePolicy* policy, + bool enable_metrics); + + void Setup(int phase, TensorMap& env); + void Del(TensorMap& env); + + void Unprep(int phase, TensorMap& env); + void Fetch(int phase, TensorMap& env); + void Update(int phase, TensorMap& env); + + // The ordinary forward: logits processing, sampling, and stop criteria. + void Forward(int phase, TensorMap& env); +}; + +} // namespace turbomind diff --git a/src/turbomind/generation/guided_decoding.cc b/src/turbomind/generation/guided_decoding.cc index 7a40bb0791..357e5951dc 100644 --- a/src/turbomind/generation/guided_decoding.cc +++ b/src/turbomind/generation/guided_decoding.cc @@ -44,7 +44,7 @@ void GuidedDecoding::Setup(int phase, TensorMap& env) d.matchers.clear(); d.active = false; for (const auto& r : rs) { - if (!r->generating) { + if (!r->submitted->generating) { continue; } if (d.matchers.emplace_back(r->req->matcher)) { diff --git a/src/turbomind/generation/logits_processor.cc b/src/turbomind/generation/logits_processor.cc index c32030f487..fa50c37173 100644 --- a/src/turbomind/generation/logits_processor.cc +++ b/src/turbomind/generation/logits_processor.cc @@ -55,11 +55,13 @@ struct LogitsProcessor::Data { bool has_temperature_penalty{}; }; -LogitsProcessor::LogitsProcessor(const BaseGenerationParam& base, int phases): BaseGenerationParam{base} +LogitsProcessor::LogitsProcessor( + const BaseGenerationParam& base, int phases, bool speculative_engine, int parameter_capacity): + BaseGenerationParam{base}, speculative_engine_{speculative_engine} { - buf_ = std::make_shared(max_batch_size_, kCPUpinned); + buf_ = std::make_shared(parameter_capacity, kCPUpinned); for (int i = 0; i < phases; ++i) { - data_.push_back(std::make_shared(max_batch_size_, kDEVICE)); + data_.push_back(std::make_shared(parameter_capacity, kDEVICE)); } } @@ -82,12 +84,12 @@ void LogitsProcessor::Forward(int phase, TensorMap& env) // repetition penalty if (d.has_repetition_penalty) { - ApplyRepetitionPenalty(logits, d.repetition_penalty_buf, token_ids_ptrs, sequence_length, stream); + ApplyRepetitionPenalty(logits, d.repetition_penalty_buf, token_ids_ptrs, sequence_length, nullptr, stream); } // ban bad words if (auto& bad_words = d.bad_words_ten) { - BanBadWords(logits, token_ids_ptrs, sequence_length, bad_words, stream); + BanBadWords(logits, token_ids_ptrs, sequence_length, bad_words, nullptr, stream); } // min length @@ -116,6 +118,49 @@ void LogitsProcessor::Forward(int phase, TensorMap& env) TM_LOG_DEBUG("{} stop", __PRETTY_FUNCTION__); } +void LogitsProcessor::ForwardVerificationBlock(int phase, + Tensor_ logits, + const Buffer_& token_ids_ptrs, + const Buffer_& effective_history, + const Buffer_& logits_active) +{ + TM_FUNCTION_SCOPE(); + TM_LOG_DEBUG("{} start", __PRETTY_FUNCTION__); + + const auto bsz = logits.shape(0); + + auto& d = *data_.at(phase); + + auto stream = core::Context::stream().handle(); + + if (d.has_repetition_penalty) { + ApplyRepetitionPenalty( + logits, d.repetition_penalty_buf, token_ids_ptrs, effective_history, logits_active.data(), stream); + } + + if (auto& bad_words = d.bad_words_ten) { + BanBadWords(logits, token_ids_ptrs, effective_history, bad_words, logits_active.data(), stream); + } + + if (d.has_min_length_penalty) { + TM_SCOPE_CALL(invokeMinLengthPenalty(logits.data(), + d.min_lengths_buf.data(), + effective_history.data(), + vocab_size_padded_, + bsz, + d.end_ids_ten.data(), + d.end_ids_ten.shape(1), + stream)); + } + + if (d.has_temperature_penalty) { + TM_SCOPE_CALL(invokeBatchApplyTemperaturePenalty_v2( + logits.data(), (float*)nullptr, d.temperature_buf.data(), bsz, vocab_size_, vocab_size_padded_, stream)); + } + + TM_LOG_DEBUG("{} stop", __PRETTY_FUNCTION__); +} + void LogitsProcessor::Setup(int phase, TensorMap& env) { TM_FUNCTION_SCOPE(); @@ -126,9 +171,31 @@ void LogitsProcessor::Setup(int phase, TensorMap& env) // const auto& rs = env.at("batch").data()[0]->rc; Buffer_ rs = env.at("requests").buffer(); + std::vector parameter_requests; + Buffer_ setup_requests = rs; + + if (speculative_engine_) { + std::vector generating_requests; + generating_requests.reserve(rs.size()); + for (Sequence* request : rs) { + if (request->submitted->generating) { + generating_requests.push_back(request); + } + } + const int P = env.at("verification_positions").data()[0]; + const int G = generating_requests.size(); + parameter_requests.resize(P * G); + for (int position = 0; position < P; ++position) { + for (int g = 0; g < G; ++g) { + parameter_requests[position * G + g] = generating_requests[g]; + } + } + setup_requests = {parameter_requests.data(), static_cast(parameter_requests.size()), kCPU}; + } + auto& copy = *env.at("copy").data()[0]; - const int bsz = rs.size(); + const int bsz = setup_requests.size(); auto& repetition_penalty = buf_->repetition_penalty_buf; auto& temperature = buf_->temperature_buf; @@ -140,7 +207,7 @@ void LogitsProcessor::Setup(int phase, TensorMap& env) d.has_bad_words_penalty = {}; for (int i = 0; i < bsz; ++i) { - auto& g = rs[i]->gen_cfg; + auto& g = setup_requests[i]->gen_cfg; // repetition_penalty repetition_penalty[i] = g.repetition_penalty; @@ -155,8 +222,8 @@ void LogitsProcessor::Setup(int phase, TensorMap& env) } // min_length - min_lengths[i] = rs[i]->prompt_len + g.min_new_tokens; - if (rs[i]->seq_len + rs[i]->inflight_new_tokens < min_lengths[i]) { + min_lengths[i] = setup_requests[i]->prompt_len + g.min_new_tokens; + if (setup_requests[i]->seq_len + setup_requests[i]->inflight_new_tokens < min_lengths[i]) { d.has_min_length_penalty = true; } } @@ -176,7 +243,7 @@ void LogitsProcessor::Setup(int phase, TensorMap& env) d.bad_words_ten = {}; init_stop_bad_words(&GenerationConfig::bad_ids, // "bad_words", - rs, + setup_requests, buf_->bad_words_buf.data(), d.bad_words_buf.data(), d.bad_words_ten, @@ -186,20 +253,20 @@ void LogitsProcessor::Setup(int phase, TensorMap& env) d.end_ids_ten = {}; int max_length = 0; for (int i = 0; i < bsz; ++i) { - max_length = std::max(max_length, (int)rs[i]->gen_cfg.eos_ids.size()); + max_length = std::max(max_length, (int)setup_requests[i]->gen_cfg.eos_ids.size()); } if (max_length) { max_length = std::min(max_length, kMaxEndIdsSize); int* h_end_ids = buf_->end_ids_buf.data(); std::fill(h_end_ids, h_end_ids + std::min(kMaxEndIdsSize, max_length) * bsz, -1); for (int i = 0; i < bsz; ++i) { - const auto& eos_ids = rs[i]->gen_cfg.eos_ids; + const auto& eos_ids = setup_requests[i]->gen_cfg.eos_ids; if (eos_ids.size() == 0) { continue; } if (TM_UNLIKELY(eos_ids.size() > kMaxEndIdsSize)) { TM_LOG_WARN("ID {}: eos length ({}) exceeds {}, truncated to {}", - rs[i]->req->id, + setup_requests[i]->req->id, eos_ids.size(), kMaxEndIdsSize, kMaxEndIdsSize); diff --git a/src/turbomind/generation/logits_processor.h b/src/turbomind/generation/logits_processor.h index 3de1293eab..4da342e043 100644 --- a/src/turbomind/generation/logits_processor.h +++ b/src/turbomind/generation/logits_processor.h @@ -26,15 +26,24 @@ namespace turbomind { class LogitsProcessor: public BaseGenerationParam { public: - explicit LogitsProcessor(const BaseGenerationParam& base, int phases); + explicit LogitsProcessor( + const BaseGenerationParam& base, int phases, bool speculative_engine, int parameter_capacity); void Setup(int phase, TensorMap& env); void Forward(int phase, TensorMap& env); + void ForwardVerificationBlock(int phase, + Tensor_ logits, + const Buffer_& token_ids_ptrs, + const Buffer_& effective_history, + const Buffer_& logits_active); + private: struct Data; + bool speculative_engine_; + std::vector> data_; std::shared_ptr buf_; // temp host buffer diff --git a/src/turbomind/generation/sampling.cc b/src/turbomind/generation/sampling.cc index ff3266761a..dcc2b06cbe 100644 --- a/src/turbomind/generation/sampling.cc +++ b/src/turbomind/generation/sampling.cc @@ -19,6 +19,7 @@ #include "src/turbomind/kernels/sampling_kernels.h" #include "src/turbomind/kernels/sampling_topk_kernels.h" #include "src/turbomind/kernels/sampling_topp_kernels.h" +#include "src/turbomind/kernels/speculative_sampling_kernels.h" #include "src/turbomind/utils/cuda_utils.h" #include "src/turbomind/engine/batch.h" @@ -37,12 +38,13 @@ struct SamplingData { std::shared_ptr request; }; - explicit SamplingData(int max_batch_size, DeviceType device) + explicit SamplingData(int parameter_capacity, int max_batch_size, DeviceType device) { - top_k_buf = {max_batch_size, device}; - top_p_buf = {max_batch_size, device}; - min_p_buf = {max_batch_size, device}; - kept_buf = {max_batch_size, device}; + top_k_buf = {parameter_capacity, device}; + top_p_buf = {parameter_capacity, device}; + min_p_buf = {parameter_capacity, device}; + kept_buf = {parameter_capacity, device}; + greedy = {parameter_capacity, device}; sampled_logprobs = {max_batch_size * (ssize_t)kMaxLogProb, device}; sampled_indices = {max_batch_size * (ssize_t)kMaxLogProb, device}; @@ -58,7 +60,8 @@ struct SamplingData { Buffer_ top_p_buf; Buffer_ min_p_buf; - Buffer_ kept_buf; // kept sample + Buffer_ kept_buf; // kept sample + Buffer_ greedy; int generation_size = 0; bool output_logprobs = 0; @@ -69,44 +72,32 @@ struct SamplingData { Buffer_ sampled_nums; }; -Sampling::Sampling(const BaseGenerationParam& base, int phases, int tp_rank): - BaseGenerationParam{base}, tp_rank_{tp_rank} +Sampling::Sampling(const BaseGenerationParam& base, int phases, int tp_rank, int parameter_capacity): + BaseGenerationParam{base}, tp_rank_{tp_rank}, parameter_capacity_{parameter_capacity} { - top_k_ = {max_batch_size_, kCPUpinned}; - top_p_ = {max_batch_size_, kCPUpinned}; - min_p_ = {max_batch_size_, kCPUpinned}; - kept_ = {max_batch_size_, kCPUpinned}; + top_k_ = {parameter_capacity_, kCPUpinned}; + top_p_ = {parameter_capacity_, kCPUpinned}; + min_p_ = {parameter_capacity_, kCPUpinned}; + kept_ = {parameter_capacity_, kCPUpinned}; + greedy_ = {parameter_capacity_, kCPUpinned}; sampled_logprobs_buf_ = {max_batch_size_ * (ssize_t)kMaxLogProb, kCPUpinned}; sampled_indices_buf_ = {max_batch_size_ * (ssize_t)kMaxLogProb, kCPUpinned}; sampled_nums_buf_ = {max_batch_size_, kCPUpinned}; // constant array - std::fill_n(kept_.data(), max_batch_size_, vocab_size_); + std::fill_n(kept_.data(), parameter_capacity_, vocab_size_); for (int i = 0; i < phases; ++i) { - data_.push_back(std::make_shared(max_batch_size_, kDEVICE)); + data_.push_back(std::make_shared(parameter_capacity_, max_batch_size_, kDEVICE)); } } -void Sampling::Forward(int phase, TensorMap& args) +void Sampling::ProcessDistributions(int phase, Tensor_ probabilities, Buffer_ token_indices) { - TM_FUNCTION_SCOPE(); - // step1: - // - use topk / topp_minp kernel to sort and filter the scores - // - softmax the left score - // step2: - // - sampling from left and sorted scores - - TM_LOG_DEBUG("{} start", __PRETTY_FUNCTION__); - auto& d = *data_.at(phase); - Tensor_ logits = args.at("logits"); - - const auto bsz = logits.shape(0); - - Buffer_ indices(bsz * vocab_size_padded_, kDEVICE); + const auto bsz = probabilities.shape(0); auto stream = core::Context::stream().handle(); @@ -114,9 +105,9 @@ void Sampling::Forward(int phase, TensorMap& args) if (d.max_topk > 0) { // TODO: top_k >= 64 is much slower than torch.topk() TopKSortFilterParams params{}; - params.logits = logits.data(); - params.sorted_logits = logits.data(); - params.sorted_indices = indices.data(); + params.logits = probabilities.data(); + params.sorted_logits = probabilities.data(); + params.sorted_indices = token_indices.data(); params.kept = d.kept_buf.data(); params.top_ks = d.top_k_buf.data(); params.max_top_k = d.max_topk; @@ -128,13 +119,13 @@ void Sampling::Forward(int phase, TensorMap& args) // use topp sort if some request skip topk filter if (d.min_topk == 0) { - TM_SCOPE_CALL( - invokeSoftmax(logits.data(), vocab_size_padded_, vocab_size_, bsz, d.kept_buf.data(), stream)); + TM_SCOPE_CALL(invokeSoftmax( + probabilities.data(), vocab_size_padded_, vocab_size_, bsz, d.kept_buf.data(), stream)); TopPSortParams params{}; - params.logits = logits.data(); - params.sorted_logits = logits.data(); - params.sorted_indices = indices.data(); + params.logits = probabilities.data(); + params.sorted_logits = probabilities.data(); + params.sorted_indices = token_indices.data(); params.kept = d.kept_buf.data(); params.top_ks = d.top_k_buf.data(); params.top_ps = d.top_p_buf.data(); @@ -147,8 +138,8 @@ void Sampling::Forward(int phase, TensorMap& args) // apply topp minp filter if (d.max_minp != 0.f || d.min_topp != 1.f) { TopPMinPFilterParams params{}; - params.sorted_logits = logits.data(); - params.sorted_indices = indices.data(); + params.sorted_logits = probabilities.data(); + params.sorted_indices = token_indices.data(); params.kept = d.kept_buf.data(); params.top_ps = d.top_p_buf.data(); params.min_ps = d.min_p_buf.data(); @@ -157,19 +148,43 @@ void Sampling::Forward(int phase, TensorMap& args) params.vocab_size_padded = vocab_size_padded_; TM_SCOPE_CALL(invokeTopPMinPFilter(params, stream)); } +} + +void Sampling::Forward(int phase, TensorMap& args) +{ + TM_FUNCTION_SCOPE(); + // step1: + // - use topk / topp_minp kernel to sort and filter the scores + // - softmax the left score + // step2: + // - sampling from left and sorted scores + + TM_LOG_DEBUG("{} start", __PRETTY_FUNCTION__); + + auto& d = *data_.at(phase); + + Tensor_ logits = args.at("logits"); + + const auto bsz = logits.shape(0); + + Buffer_ indices(bsz * vocab_size_padded_, kDEVICE); + + auto stream = core::Context::stream().handle(); + + ProcessDistributions(phase, logits, indices); // sample { SamplingParams params{}; - params.logits = logits.data(); + params.probabilities = logits.data(); params.stride = vocab_size_padded_; params.indices = indices.data(); params.kept = d.kept_buf.data(); params.curandstate = (curandState_t*)args.at("curand_state").raw_data(); params.curandstate_indices = args.at("curand_state_indices").data(); + params.sample_mask = args.contains("sample_mask") ? args.at("sample_mask").data() : nullptr; params.batch_size = bsz; - params.output_ids = args.at("output_ids").data(); // (B, 1) - params.sequence_length = args.at("sequence_length").data(); + params.selected_tokens = args.at("output_ids").data(); // (B, 1) if (d.output_logprobs) { params.sampled_logprobs = d.sampled_logprobs.data(); @@ -183,6 +198,26 @@ void Sampling::Forward(int phase, TensorMap& args) TM_LOG_DEBUG("{} stop", __PRETTY_FUNCTION__); } +void Sampling::VerifyTargetBlock(int phase, Tensor_ probabilities, VerifyTargetBlockParams params) +{ + auto& d = *data_.at(phase); + + const int rows = probabilities.shape(0); + + Buffer_ token_indices(rows * (ssize_t)vocab_size_padded_, kDEVICE); + + ProcessDistributions(phase, probabilities, token_indices); + + params.probabilities = probabilities.data(); + params.probability_stride = probabilities.stride(0); + params.probability_token_ids = token_indices.data(); + params.token_id_stride = probabilities.stride(0); + params.kept_count = d.kept_buf.data(); + params.greedy = d.greedy.data(); + + invokeVerifyTargetBlock(params, core::Context::stream().handle()); +} + void Sampling::Setup(int phase, TensorMap& env) { TM_FUNCTION_SCOPE(); @@ -198,17 +233,16 @@ void Sampling::Setup(int phase, TensorMap& env) d.output_logprobs = false; d.logprob_outputs.clear(); - for (int i = 0; i < rc.size(); ++i) { - auto& c = *rc[i]; - if (!c.generating) { + std::vector generating_requests; + generating_requests.reserve(rc.size()); + for (Sequence* request : rc) { + auto& c = *request; + if (!c.submitted->generating) { continue; } const int row = d.generation_size++; - - top_k_[row] = c.gen_cfg.top_k; - top_p_[row] = c.gen_cfg.top_p; - min_p_[row] = c.gen_cfg.min_p; + generating_requests.push_back(request); if (c.gen_cfg.output_logprobs) { d.output_logprobs = true; @@ -216,7 +250,9 @@ void Sampling::Setup(int phase, TensorMap& env) } } - const int bsz = d.generation_size; + const int G = d.generation_size; + const int P = parameter_capacity_ == max_batch_size_ ? 1 : env.at("verification_positions").data()[0]; + const int bsz = P * G; if (bsz == 0) { d.max_topk = d.min_topk = 0; d.min_topp = 0.f; @@ -224,6 +260,17 @@ void Sampling::Setup(int phase, TensorMap& env) return; } + for (int position = 0; position < P; ++position) { + for (int g = 0; g < G; ++g) { + const int row = position * G + g; + const auto& config = generating_requests[g]->gen_cfg; + top_k_[row] = config.top_k; + top_p_[row] = config.top_p; + min_p_[row] = config.min_p; + greedy_[row] = config.top_k == 1; + } + } + d.max_topk = *std::max_element(top_k_.begin(), top_k_.begin() + bsz); d.min_topk = *std::min_element(top_k_.begin(), top_k_.begin() + bsz); d.min_topp = *std::min_element(top_p_.begin(), top_p_.begin() + bsz); @@ -234,6 +281,7 @@ void Sampling::Setup(int phase, TensorMap& env) copy(min_p_.data(), bsz, d.min_p_buf.data()); copy(kept_.data(), bsz, d.kept_buf.data()); + copy(greedy_.data(), bsz, d.greedy.data()); } void Sampling::Fetch(int phase, TensorMap& env) diff --git a/src/turbomind/generation/sampling.h b/src/turbomind/generation/sampling.h index 8cd77e18aa..f45bb4433b 100644 --- a/src/turbomind/generation/sampling.h +++ b/src/turbomind/generation/sampling.h @@ -3,6 +3,7 @@ #include "src/turbomind/core/core.h" #include "src/turbomind/generation/base_param.h" +#include "src/turbomind/kernels/speculative_sampling_kernels.h" namespace turbomind { @@ -10,18 +11,23 @@ struct SamplingData; class Sampling: public BaseGenerationParam { public: - explicit Sampling(const BaseGenerationParam& base, int phases, int tp_rank); + explicit Sampling(const BaseGenerationParam& base, int phases, int tp_rank, int parameter_capacity); void Setup(int phase, TensorMap& env); void Forward(int phase, TensorMap& env); + void VerifyTargetBlock(int phase, Tensor_ probabilities, VerifyTargetBlockParams params); + void Fetch(int phase, TensorMap& env); void Update(int phase, TensorMap& env); private: + void ProcessDistributions(int phase, Tensor_ probabilities, Buffer_ token_indices); + const int tp_rank_; + const int parameter_capacity_; std::vector> data_; @@ -30,6 +36,7 @@ class Sampling: public BaseGenerationParam { Buffer_ top_k_; Buffer_ top_p_; Buffer_ min_p_; + Buffer_ greedy_; Buffer_ sampled_logprobs_buf_; Buffer_ sampled_indices_buf_; diff --git a/src/turbomind/generation/stop_criteria.cc b/src/turbomind/generation/stop_criteria.cc index 98f73ce4ab..d57bcd65f2 100644 --- a/src/turbomind/generation/stop_criteria.cc +++ b/src/turbomind/generation/stop_criteria.cc @@ -69,23 +69,56 @@ void StopCriteria::Forward(int phase, TensorMap& env) auto stream = core::Context::stream().handle(); - if (auto& stop_words = d.stop_words_ten) { - TM_CHECK_EQ(stop_words.ndim(), 3); // [batch, 2, len] - size_t stop_words_len = stop_words.shape(2); - TM_SCOPE_CALL(invokeStopWordsCriterion_v2((const int**)token_ids_ptrs.data(), - sequence_length.data(), - stop_words.data(), - finished.data(), - stop_words_len, - batch_size, - stream)); + const int* stop_words = nullptr; + int stop_words_width = 0; + + if (d.stop_words_ten) { + stop_words = d.stop_words_ten.data(); + stop_words_width = static_cast(d.stop_words_ten.shape(2)); + } + + TM_SCOPE_CALL(invokeStopCriteria(reinterpret_cast(token_ids_ptrs.data()), + sequence_length.data(), + stop_words, + stop_words_width, + d.max_seq_len.data(), + finished.data(), + batch_size, + stream)); +} + +void StopCriteria::ForwardSpeculative(int phase, + const Buffer_& token_ids_ptrs, + const Buffer_& entry_sequence_length, + Buffer_ accept_len, + Buffer_ finished, + TensorMap& env) +{ + TM_FUNCTION_SCOPE(); + auto& d = *data_.at(phase); + + const int batch_size = token_ids_ptrs.size(); + auto stream = core::Context::stream().handle(); + + const int* stop_words = nullptr; + int stop_words_width = 0; + + if (d.stop_words_ten) { + stop_words = d.stop_words_ten.data(); + stop_words_width = static_cast(d.stop_words_ten.shape(2)); } - TM_SCOPE_CALL(invokeLengthCriterion_v2(finished.data(), // - sequence_length.data(), - d.max_seq_len.data(), - batch_size, - stream)); + static_cast(env); + + invokeStopCriteria(reinterpret_cast(token_ids_ptrs.data()), + entry_sequence_length.data(), + accept_len.data(), + stop_words, + stop_words_width, + d.max_seq_len.data(), + finished.data(), + batch_size, + stream); } } // namespace turbomind diff --git a/src/turbomind/generation/stop_criteria.h b/src/turbomind/generation/stop_criteria.h index 7daeb113ca..887234740e 100644 --- a/src/turbomind/generation/stop_criteria.h +++ b/src/turbomind/generation/stop_criteria.h @@ -32,6 +32,13 @@ class StopCriteria: public BaseGenerationParam { void Forward(int phase, TensorMap& env); + void ForwardSpeculative(int phase, + const Buffer_& token_ids_ptrs, + const Buffer_& entry_sequence_length, + Buffer_ accept_len, + Buffer_ finished, + TensorMap& env); + private: std::vector> data_; diff --git a/src/turbomind/generation/target_verification.cc b/src/turbomind/generation/target_verification.cc new file mode 100644 index 0000000000..3561f2bbbb --- /dev/null +++ b/src/turbomind/generation/target_verification.cc @@ -0,0 +1,245 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#include "src/turbomind/generation/generation_impl.h" + +#include "src/turbomind/core/allocator.h" +#include "src/turbomind/core/check.h" +#include "src/turbomind/core/copy.h" +#include "src/turbomind/engine/batch.h" +#include "src/turbomind/engine/request.h" + +#include "src/turbomind/generation/logits_processor.h" +#include "src/turbomind/generation/sampling.h" +#include "src/turbomind/generation/stop_criteria.h" + +#include "src/turbomind/kernels/sampling_topk_kernels.h" // InitializeRandomStates +#include "src/turbomind/kernels/speculative_sequence_kernels.h" + +#include "src/turbomind/models/llama/llama_kernels.h" +#include "src/turbomind/models/speculative/speculative_model.h" +#include "src/turbomind/utils/cuda_utils.h" + +namespace turbomind { + +TargetVerification::TargetVerification(GenerationShared& shared, + const SpeculativePolicy& policy, + int hidden_units, + DataType hidden_dtype, + bool enable_metrics, + int phases): + shared_{shared}, + draft_count_{policy.max_proposals()}, + enable_metrics_{enable_metrics} +{ + const int max_batch_size = shared_.max_batch_size; + const int K = draft_count_ + 1; + + request_token_ids_ptrs_buf_ = {max_batch_size, kCPUpinned}; + speculative_row_buf_ = {max_batch_size, kCPUpinned}; + selected_span_ids_buf_ = {max_batch_size * (ssize_t)K, kCPUpinned}; + accept_len_buf_ = {max_batch_size, kCPUpinned}; + if (enable_metrics_) { + accepted_draft_count_buf_ = {max_batch_size, kCPUpinned}; + } + + for (int i = 0; i < phases; ++i) { + auto d = std::make_unique(); + + d->request_token_ids_ptrs = {max_batch_size, kDEVICE}; + d->selected_span_ids = {max_batch_size * (ssize_t)K, kDEVICE}; + d->accept_len = {max_batch_size, kDEVICE}; + d->finished_on_entry = {max_batch_size, kDEVICE}; + d->speculative_row = {max_batch_size, kDEVICE}; + + if (enable_metrics_ && draft_count_ > 0) { + d->accepted_draft_count = {max_batch_size, kDEVICE}; + } + + d->block_logits_active = {K * (ssize_t)max_batch_size, kDEVICE}; + d->effective_history = {K * (ssize_t)max_batch_size, kDEVICE}; + d->verification_draft_ids = {draft_count_ * (ssize_t)max_batch_size, kDEVICE}; + d->selected_hidden = {{max_batch_size * (ssize_t)K, hidden_units}, hidden_dtype, kDEVICE}; + + data_.push_back(std::move(d)); + } +} + +TargetVerification::~TargetVerification() = default; + +void TargetVerification::Setup(int phase, TensorMap& env) +{ + auto& copy = *env.at("copy").data()[0]; + auto& d = *data_.at(phase); + auto& g = *shared_.data.at(phase); + + Buffer_ rc = env.at("requests").buffer(); + + for (int i = 0; i < rc.size(); ++i) { + auto& c = *rc[i]; + + request_token_ids_ptrs_buf_[i] = + c.generation_token_ids_row >= 0 ? shared_.RowPtr(c.generation_token_ids_row) : nullptr; + speculative_row_buf_[i] = c.submitted->is_verification_row(); + } + + const int position_count = env.at("verification_positions").data()[0]; + for (int position = 1; position < position_count; ++position) { + std::copy_n(shared_.token_ids_ptrs_buf.data(), + g.generation_size, + shared_.token_ids_ptrs_buf.data() + position * g.generation_size); + } + copy(shared_.token_ids_ptrs_buf, position_count * g.generation_size, g.token_ids_ptrs); + copy(request_token_ids_ptrs_buf_, rc.size(), d.request_token_ids_ptrs); + copy(speculative_row_buf_, rc.size(), d.speculative_row); +} + +void TargetVerification::PublishDraftInputs(int phase, TensorMap& env) +{ + auto& d = *data_.at(phase); + auto& g = *shared_.data.at(phase); + + env.produce("request_token_ids_ptrs", d.request_token_ids_ptrs.slice(0, g.request_count)); + env.produce("request_to_generation_row_offsets", + g.request_to_generation_row_offsets.slice(0, g.request_count + 1)); + env.produce("accept_len", d.accept_len.slice(0, g.request_count)); + env.produce("finished_on_entry", d.finished_on_entry.slice(0, g.request_count)); + env.produce("speculative_row", d.speculative_row.slice(0, g.request_count)); +} + +void TargetVerification::Fetch(int phase, TensorMap& env) +{ + auto& d = *data_.at(phase); + auto& copy = *env.at("copy").data()[0]; + + auto& batch = *env.at("batch").data()[0]; + const int B = batch.bsz; + const int K = draft_count_ + 1; + + copy(d.selected_span_ids, B * K, selected_span_ids_buf_); + copy(d.accept_len, B, accept_len_buf_); + + env.produce("selected_span_ids", selected_span_ids_buf_.slice(0, B * K)); + env.produce("accept_len", accept_len_buf_.slice(0, B)); + + if (enable_metrics_) { + copy(d.accepted_draft_count, B, accepted_draft_count_buf_); + env.produce("accepted_draft_count", accepted_draft_count_buf_.slice(0, B)); + } +} + +Tensor TargetVerification::SelectedHiddenBuffer(int phase, core::ssize_t rows) +{ + return data_.at(phase)->selected_hidden.slice(0, rows); +} + +void TargetVerification::InitializeTargetVerification(int phase, int position_count, TensorMap& env) +{ + TM_FUNCTION_SCOPE(); + auto& d = *data_.at(phase); + auto& g = *shared_.data.at(phase); + + const auto stream = core::Context::stream().handle(); + + Copy(env.at("finished").buffer(), g.request_count, d.finished_on_entry); + Clear(d.accept_len.slice(0, g.request_count)); + + if (g.random_init_needed) { + InitializeRandomStates((curandState_t*)shared_.random_states.raw_data(), + g.random_seed.data(), + g.random_init.data(), + shared_.max_batch_size, + stream); + } + + const Buffer_ entry_sequence_length = env.at("sequence_length").buffer(); + + invokeInitializeTargetVerification(d.block_logits_active.data(), + d.effective_history.data(), + d.verification_draft_ids.data(), + reinterpret_cast(d.request_token_ids_ptrs.data()), + entry_sequence_length.data(), + d.finished_on_entry.data(), + d.speculative_row.data(), + enable_metrics_ ? d.accepted_draft_count.data() : nullptr, + g.request_to_generation_row_offsets.data(), + g.request_count, + g.generation_size, + position_count, + stream); +} + +void TargetVerification::ProcessTargetBlock(int phase, int position_count, const Tensor& target_logits, TensorMap& env) +{ + TM_FUNCTION_SCOPE(); + auto& d = *data_.at(phase); + auto& g = *shared_.data.at(phase); + + const int B = g.request_count; + const int G = g.generation_size; + const int rows = position_count * G; + + const auto stream = core::Context::stream().handle(); + + if (rows == 0) { + return; + } + + Tensor_ probabilities = empty_like(target_logits, kFloat32); + invokeCastFloat2D(target_logits, probabilities, stream); + + shared_.logits_processor->ForwardVerificationBlock(phase, + probabilities, + g.token_ids_ptrs.slice(0, rows), + d.effective_history.slice(0, rows), + d.block_logits_active.slice(0, rows)); + + VerifyTargetBlockParams p{}; + p.verification_draft_ids = d.verification_draft_ids.data(); + p.draft_row_stride = G; + p.logits_active = d.block_logits_active.data(); + p.random_states = reinterpret_cast(shared_.random_states.raw_data()); + p.random_state_indices = g.random_state_indices.data(); + p.request_token_ids_ptrs = d.request_token_ids_ptrs.data(); + p.entry_sequence_length = env.at("sequence_length").data(); + p.request_to_generation_offsets = g.request_to_generation_row_offsets.data(); + p.speculative_row = d.speculative_row.data(); + p.selected_span_ids = d.selected_span_ids.data(); + p.selected_span_stride = draft_count_ + 1; + p.accept_len = d.accept_len.data(); + p.accepted_draft_count = enable_metrics_ ? d.accepted_draft_count.data() : nullptr; + p.request_count = B; + p.generation_count = G; + p.position_count = position_count; + + shared_.sampling->VerifyTargetBlock(phase, probabilities, p); +} + +void TargetVerification::ClampSelectedSpan(int phase, TensorMap& env) +{ + TM_FUNCTION_SCOPE(); + auto& d = *data_.at(phase); + auto& g = *shared_.data.at(phase); + + const Buffer_ entry_sequence_length = env.at("sequence_length").buffer(); + + shared_.stop_criteria->ForwardSpeculative(phase, + d.request_token_ids_ptrs.slice(0, g.request_count), + entry_sequence_length, + d.accept_len.slice(0, g.request_count), + env.at("finished").buffer(), + env); +} + +void TargetVerification::CommitAcceptedSpan(int phase, Buffer_ sequence_length) +{ + auto& d = *data_.at(phase); + auto& g = *shared_.data.at(phase); + + invokeAdvanceSequenceByAcceptedSpan(sequence_length.data(), + d.accept_len.data(), + enable_metrics_ ? d.accepted_draft_count.data() : nullptr, + g.request_count, + core::Context::stream().handle()); +} + +} // namespace turbomind diff --git a/src/turbomind/generation/target_verification.h b/src/turbomind/generation/target_verification.h new file mode 100644 index 0000000000..5733ea537f --- /dev/null +++ b/src/turbomind/generation/target_verification.h @@ -0,0 +1,82 @@ +// Copyright (c) OpenMMLab. All rights reserved. +#pragma once + +#include +#include + +#include "src/turbomind/core/core.h" +#include "src/turbomind/engine/batch.h" + +namespace turbomind { + +class SpeculativePolicy; +struct GenerationShared; + +// Spec-side verification half of a speculative round: verifier initialization, +// target-block verification processing, stop-span clamping, and accepted-state +// commit, together with their per-phase buffers and staging. Constructed by +// the generation module only when a speculative policy is present and driven +// by the composed executor's speculative-round routine; it borrows the +// generation module's rows, random states, and sampler infrastructure through +// GenerationShared. +class TargetVerification { +public: + TargetVerification(GenerationShared& shared, + const SpeculativePolicy& policy, + int hidden_units, + DataType hidden_dtype, + bool enable_metrics, + int phases); + ~TargetVerification(); + + // Spec-side participation in the shared BatchOp lifecycle. + void Setup(int phase, TensorMap& env); + void Fetch(int phase, TensorMap& env); + + // Publishes the verification draft inputs the composed round reads: + // request token-row pointers, request-to-generation row offsets, accepted + // lengths, entry finished flags, and the per-row speculative flag. Driven + // by the executor's prepare bracket after generation's prepare and before + // the draft's, so the draft's dependency is bracket code, not fanout + // order. + void PublishDraftInputs(int phase, TensorMap& env); + + // The round's selected-states storage, sliced to the selected row count. + // The executor publishes it as the target decoder's selected-hidden + // buffer before the target pass. + Tensor SelectedHiddenBuffer(int phase, core::ssize_t rows); + + // Driven by the speculative-round routine. + void InitializeTargetVerification(int phase, int position_count, TensorMap& env); + void ProcessTargetBlock(int phase, int position_count, const Tensor& target_logits, TensorMap& env); + void ClampSelectedSpan(int phase, TensorMap& env); + void CommitAcceptedSpan(int phase, Buffer_ sequence_length); + +private: + struct Data { + Buffer_ request_token_ids_ptrs; + Buffer_ selected_span_ids; + Buffer_ accept_len; + Buffer_ finished_on_entry; + Buffer_ accepted_draft_count; + Buffer_ speculative_row; + Buffer_ block_logits_active; + Buffer_ effective_history; + Buffer_ verification_draft_ids; + Tensor selected_hidden; + }; + + GenerationShared& shared_; + const int draft_count_; + const bool enable_metrics_; + + std::vector> data_; + + Buffer_ request_token_ids_ptrs_buf_; + Buffer_ speculative_row_buf_; + Buffer_ selected_span_ids_buf_; + Buffer_ accept_len_buf_; + Buffer_ accepted_draft_count_buf_; +}; + +} // namespace turbomind diff --git a/src/turbomind/kernels/CMakeLists.txt b/src/turbomind/kernels/CMakeLists.txt index 03d996388f..c50edef0cf 100644 --- a/src/turbomind/kernels/CMakeLists.txt +++ b/src/turbomind/kernels/CMakeLists.txt @@ -77,6 +77,34 @@ add_library(sampling_kernels STATIC sampling_kernels.cu) set_property(TARGET sampling_kernels PROPERTY POSITION_INDEPENDENT_CODE ON) set_property(TARGET sampling_kernels PROPERTY CUDA_RESOLVE_DEVICE_SYMBOLS ON) +add_library(speculative_sampling_kernels STATIC + speculative_sampling_kernels.cu) +set_property(TARGET speculative_sampling_kernels + PROPERTY POSITION_INDEPENDENT_CODE ON) +set_property(TARGET speculative_sampling_kernels + PROPERTY CUDA_RESOLVE_DEVICE_SYMBOLS ON) +target_link_libraries(speculative_sampling_kernels PRIVATE + core + CUDA::curand) + +add_library(draft_carry_kernels STATIC + draft_carry_kernels.cu) +set_property(TARGET draft_carry_kernels + PROPERTY POSITION_INDEPENDENT_CODE ON) +set_property(TARGET draft_carry_kernels + PROPERTY CUDA_RESOLVE_DEVICE_SYMBOLS ON) +target_link_libraries(draft_carry_kernels PRIVATE + core) + +add_library(speculative_sequence_kernels STATIC + speculative_sequence_kernels.cu) +set_property(TARGET speculative_sequence_kernels + PROPERTY POSITION_INDEPENDENT_CODE ON) +set_property(TARGET speculative_sequence_kernels + PROPERTY CUDA_RESOLVE_DEVICE_SYMBOLS ON) +target_link_libraries(speculative_sequence_kernels PRIVATE + core) + add_library(apply_token_bitmask_inplace_cuda STATIC apply_token_bitmask_inplace_cuda.cu) set_property(TARGET apply_token_bitmask_inplace_cuda PROPERTY POSITION_INDEPENDENT_CODE ON) set_property(TARGET apply_token_bitmask_inplace_cuda PROPERTY CUDA_RESOLVE_DEVICE_SYMBOLS ON) diff --git a/src/turbomind/kernels/attention/CMakeLists.txt b/src/turbomind/kernels/attention/CMakeLists.txt index 036ac5d529..3ea7d0b79a 100644 --- a/src/turbomind/kernels/attention/CMakeLists.txt +++ b/src/turbomind/kernels/attention/CMakeLists.txt @@ -1,6 +1,7 @@ # Copyright (c) OpenMMLab. All rights reserved. add_subdirectory(kernel) +add_subdirectory(verification) add_library(attention STATIC attention.cu @@ -14,6 +15,7 @@ set_property(TARGET attention PROPERTY CUDA_RESOLVE_DEVICE_SYMBOLS ON) target_compile_options(attention PRIVATE -O3 $<$:-use_fast_math --expt-relaxed-constexpr -Xptxas=-v --threads=${NVCC_THREADS}>) target_link_libraries(attention PUBLIC $) +target_link_libraries(attention PUBLIC verification_attention) target_link_libraries(attention PRIVATE nvidia::cutlass::cutlass) if (BUILD_TEST) diff --git a/src/turbomind/kernels/attention/cp_utils.cu b/src/turbomind/kernels/attention/cp_utils.cu index 3d5254c2b9..8a3571c368 100644 --- a/src/turbomind/kernels/attention/cp_utils.cu +++ b/src/turbomind/kernels/attention/cp_utils.cu @@ -1,5 +1,7 @@ // Copyright (c) OpenMMLab. All rights reserved. +#include + #include "src/turbomind/kernels/attention/cp_utils.h" namespace turbomind { @@ -49,4 +51,24 @@ void CpPost(void* context) TM_CUDA_CHECK(cudaStreamWaitEvent(ctx->stream, ctx->consume_event, 0)); } +__global__ void FillNegInfMLKernel(float4* data, size_t n_quads) +{ + const size_t idx = (size_t)blockIdx.x * blockDim.x + threadIdx.x; + if (idx < n_quads) { + data[idx] = make_float4(-CUDART_INF_F, 0.f, -CUDART_INF_F, 0.f); + } +} + +void invokeFillNegInfML(float* data, size_t n_pairs, cudaStream_t stream) +{ + if (n_pairs == 0) { + return; + } + constexpr int block = 256; + const size_t n_quads = n_pairs >> 1; + const size_t grid = (n_quads + block - 1) / block; + FillNegInfMLKernel<<>>(reinterpret_cast(data), n_quads); + TM_CUDA_CHECK(cudaGetLastError()); +} + } // namespace turbomind diff --git a/src/turbomind/kernels/attention/cp_utils.h b/src/turbomind/kernels/attention/cp_utils.h index 251c6ec4c2..73214dc0f8 100644 --- a/src/turbomind/kernels/attention/cp_utils.h +++ b/src/turbomind/kernels/attention/cp_utils.h @@ -33,4 +33,6 @@ struct CpPostContext { void CpPost(void* context); +void invokeFillNegInfML(float* data, size_t n_pairs, cudaStream_t stream); + } // namespace turbomind diff --git a/src/turbomind/kernels/attention/kv_cache_utils_v2.cu b/src/turbomind/kernels/attention/kv_cache_utils_v2.cu index 2e384c318d..562a529f28 100644 --- a/src/turbomind/kernels/attention/kv_cache_utils_v2.cu +++ b/src/turbomind/kernels/attention/kv_cache_utils_v2.cu @@ -25,6 +25,7 @@ __global__ void __launch_bounds__(128) ProcessKV_v2(char** blocks, const int* cu_k_len, const int* cu_block_num, const int* readonly_block_num, + const bool* finished, RopeKernelParam rope_param, int64_t stride_b, int64_t stride_c, @@ -51,6 +52,10 @@ __global__ void __launch_bounds__(128) ProcessKV_v2(char** blocks, const int head_idx = blockIdx.y; const int batch_idx = blockIdx.z; + if (finished && finished[batch_idx]) { + return; + } + const int qi_beg = cu_q_len[batch_idx]; const int qi_end = cu_q_len[batch_idx + 1]; const int q_len = qi_end - qi_beg; @@ -219,6 +224,7 @@ void invokeProcessKV_v2(char** blocks, const int* cu_k_len, const int* cu_block_num, const int* readonly_block_num, + const bool* finished, const RopeKernelParam& rope_param, int64_t stride_b, int64_t stride_c, @@ -261,6 +267,7 @@ void invokeProcessKV_v2(char** blocks, cu_k_len, cu_block_num, readonly_block_num, + finished, rope_param, stride_b, stride_c, @@ -316,6 +323,7 @@ void invokeProcessKV_v2(char** blocks, const int* cu_k_len, \ const int* cu_block_num, \ const int* readonly_block_num, \ + const bool* finished, \ const RopeKernelParam& rope_param, \ int64_t stride_b, \ int64_t stride_c, \ @@ -343,6 +351,7 @@ __global__ void __launch_bounds__(128) flattenKV_v2(T* k, const Tkv** blocks, const int* cu_k_len, const int* cu_block_num, + const bool* finished, RopeKernelParam rope_param, int64_t stride_b, int64_t stride_c, @@ -366,6 +375,10 @@ __global__ void __launch_bounds__(128) flattenKV_v2(T* k, const int head_idx = blockIdx.y; const int batch_idx = blockIdx.z; + if (finished && finished[batch_idx]) { + return; + } + const int ti_0 = cu_k_len[0]; const int ti_beg = cu_k_len[batch_idx] - ti_0; const int ti_end = cu_k_len[batch_idx + 1] - ti_0; @@ -473,6 +486,7 @@ void invokeFlattenKV_v2(T* k, char** blocks, const int* cu_k_len, const int* cu_block_num, + const bool* finished, const RopeKernelParam& rope_param, int64_t stride_b, int64_t stride_c, @@ -511,6 +525,7 @@ void invokeFlattenKV_v2(T* k, (const Tkv**)blocks, cu_k_len, cu_block_num, + finished, rope_param, stride_b, stride_c, @@ -561,6 +576,7 @@ void invokeFlattenKV_v2(T* k, char** blocks, \ const int* cu_k_len, \ const int* cu_block_num, \ + const bool* finished, \ const RopeKernelParam& rope_param, \ int64_t stride_b, \ int64_t stride_c, \ diff --git a/src/turbomind/kernels/attention/kv_cache_utils_v2.h b/src/turbomind/kernels/attention/kv_cache_utils_v2.h index b5c2df8570..cddf44a612 100644 --- a/src/turbomind/kernels/attention/kv_cache_utils_v2.h +++ b/src/turbomind/kernels/attention/kv_cache_utils_v2.h @@ -17,6 +17,7 @@ void invokeProcessKV_v2(char** blocks, const int* cu_k_len, const int* cu_block_num, const int* readonly_block_num, + const bool* finished, const RopeKernelParam& rope_param, int64_t stride_b, int64_t stride_c, @@ -45,6 +46,7 @@ void invokeProcessKV_v2_(const AttentionParams& params) params.cu_k_len, params.block_iter_params.cu_block_nums, params.readonly_block_num, + params.finished, params.rope_param, 0, // stride b params.stride / params.size_per_head, // stride c @@ -68,6 +70,7 @@ void invokeFlattenKV_v2(T* k, char** blocks, const int* cu_k_len, const int* cu_block_num, + const bool* finished, const RopeKernelParam& rope_param, int64_t stride_b, int64_t stride_c, @@ -94,6 +97,7 @@ void invokeFlattenKV_v2_(const AttentionParams& params, int sum_k_len) (char**)params.block_iter_params.block_ptrs, params.cu_k_len, params.block_iter_params.cu_block_nums, + params.finished, RopeKernelParam{}, 0, 1, diff --git a/src/turbomind/kernels/attention/rotary_embedding.h b/src/turbomind/kernels/attention/rotary_embedding.h index d4281333a0..d4d1c56e6a 100644 --- a/src/turbomind/kernels/attention/rotary_embedding.h +++ b/src/turbomind/kernels/attention/rotary_embedding.h @@ -82,8 +82,15 @@ struct FastRoPE { } // mrope is an operation applied on top of any base rope type if (param_.mrope_mode != MropeMode::kNone) { + if (param_.mrope.position_offsets) { + // Flat [rows, 3] table: per-token rows are reached through offsets. + param_.mrope.position_offsets += batch_idx; + } + else { + // Per-batch [batch, rows, 3] tables: reach this batch's table by stride. + param_.mrope.position_ids += batch_idx * param_.mrope.stride; + } param_.mrope.position_delta += batch_idx; - param_.mrope.position_offsets += batch_idx; param_.mrope.length += batch_idx; } } @@ -127,11 +134,95 @@ struct FastRoPE { } } + template + __device__ void apply(Array& x, float timestep) + { + if (param_.mrope_mode == MropeMode::kNone) { + // Most models apply rotary embedding in half precision + PRAGMA_UNROLL + for (int i = 0; i < N; i += 2) { + rotate_pair(x, i, timestep); + } + } + else if (param_.mrope_mode == MropeMode::kChunked) { + apply_mrope_impl(x, timestep); + } + else if (param_.mrope_mode == MropeMode::kInterleaved) { + apply_mrope_impl(x, timestep); + } + } + __device__ __forceinline__ MropeCoord get_mrope_coord(float timestep, int token_idx) const { - if (token_idx < *param_.mrope.length) { - const int row = *param_.mrope.position_offsets + token_idx; - const int* t = param_.mrope.position_ids + 3 * row; + if (param_.mrope.position_offsets) { + // Flat [rows, 3] table: rows are reached through per-token offsets. + if (token_idx < *param_.mrope.length) { + const int row = *param_.mrope.position_offsets + token_idx; + const int* t = param_.mrope.position_ids + 3 * row; + return {t[0], t[1], t[2]}; + } + } + else if (timestep < *param_.mrope.length) { + // Per-batch [batch, rows, 3] tables (ctor advanced to this batch's table). + const int* t = param_.mrope.position_ids + 3 * (int)timestep; + return {t[0], t[1], t[2]}; + } + const int pos = (int)timestep + (*param_.mrope.position_delta); + return {pos, pos, pos}; + } + + template + __device__ void apply(Array& x, Array& y, float timestep) + { + if (param_.mrope_mode == MropeMode::kNone) { + PRAGMA_UNROLL + for (int i = 0; i < N; i += 2) { + rotate_pair(x, y, i, timestep); + } + } + else if (param_.mrope_mode == MropeMode::kChunked) { + apply_mrope_impl(x, y, timestep); + } + else if (param_.mrope_mode == MropeMode::kInterleaved) { + apply_mrope_impl(x, y, timestep); + } + } + + template + __device__ void fill_coefficients(Array& cs, float timestep) const + { + if (param_.mrope_mode == MropeMode::kNone) { + PRAGMA_UNROLL + for (int i = 0; i < N; i += 2) { + fill_coefficient_pair(cs, i, timestep); + } + } + else if (param_.mrope_mode == MropeMode::kChunked) { + fill_mrope_coefficients(cs, timestep); + } + else if (param_.mrope_mode == MropeMode::kInterleaved) { + fill_mrope_coefficients(cs, timestep); + } + } + + template + __device__ void apply_coefficients(Array& x, const Array& cs) const + { + PRAGMA_UNROLL + for (int i = 0; i < N; i += 2) { + T tmp0 = cs[i] * x[i] - cs[i + 1] * x[i + 1]; + T tmp1 = cs[i] * x[i + 1] + cs[i + 1] * x[i]; + if (is_valid_) { + x[i] = tmp0; + x[i + 1] = tmp1; + } + } + } + + __device__ __forceinline__ MropeCoord get_mrope_coord(float timestep) const + { + if (timestep < *param_.mrope.length) { + const int* t = param_.mrope.position_ids + 3 * (int)timestep; return {t[0], t[1], t[2]}; } const int pos = (int)timestep + (*param_.mrope.position_delta); @@ -180,6 +271,38 @@ struct FastRoPE { } } + template + __device__ __forceinline__ void rotate_pair( + Array& x, Array& y, int i, float timestep) const + { + float c, s; + sincosf(timestep * inv_freq_[i / 2], &s, &c); + s *= attention_scaling_; + c *= attention_scaling_; + T x0 = (T)c * x[i] - (T)s * x[i + 1]; + T x1 = (T)c * x[i + 1] + (T)s * x[i]; + T y0 = (T)c * y[i] - (T)s * y[i + 1]; + T y1 = (T)c * y[i + 1] + (T)s * y[i]; + if (is_valid_) { + x[i] = x0; + x[i + 1] = x1; + y[i] = y0; + y[i + 1] = y1; + } + } + + template + __device__ __forceinline__ void + fill_coefficient_pair(Array& cs, int i, float timestep) const + { + float c, s; + sincosf(timestep * inv_freq_[i / 2], &s, &c); + s *= attention_scaling_; + c *= attention_scaling_; + cs[i] = (T)c; + cs[i + 1] = (T)s; + } + template __device__ __forceinline__ void apply_mrope_impl(Array& x, float timestep, int token_idx) const { @@ -191,6 +314,45 @@ struct FastRoPE { rotate_pair(x, i, (float)ts); } } + + template + __device__ __forceinline__ void apply_mrope_impl(Array& x, float timestep) const + { + const MropeCoord coord = get_mrope_coord(timestep); + PRAGMA_UNROLL + for (int i = 0; i < N; i += 2) { + const int pair_idx = (i + idx_) >> 1; + const int ts = select_mrope_timestep(pair_idx, coord); + rotate_pair(x, i, (float)ts); + } + } + + + template + __device__ __forceinline__ void apply_mrope_impl( + Array& x, Array& y, float timestep) const + { + const MropeCoord coord = get_mrope_coord(timestep); + PRAGMA_UNROLL + for (int i = 0; i < N; i += 2) { + const int pair_idx = (i + idx_) >> 1; + const int ts = select_mrope_timestep(pair_idx, coord); + rotate_pair(x, y, i, (float)ts); + } + } + + template + __device__ __forceinline__ void + fill_mrope_coefficients(Array& cs, float timestep) const + { + const MropeCoord coord = get_mrope_coord(timestep); + PRAGMA_UNROLL + for (int i = 0; i < N; i += 2) { + const int pair_idx = (i + idx_) >> 1; + const int ts = select_mrope_timestep(pair_idx, coord); + fill_coefficient_pair(cs, i, (float)ts); + } + } }; template diff --git a/src/turbomind/kernels/attention/test_attention.cu b/src/turbomind/kernels/attention/test_attention.cu index 5e06b72281..398856ab3c 100644 --- a/src/turbomind/kernels/attention/test_attention.cu +++ b/src/turbomind/kernels/attention/test_attention.cu @@ -153,6 +153,7 @@ void TestBlocks(const thrust::universal_vector& k_cache, // [B, H, S, cu_seq_lens.data().get(), cu_block_cnts.data().get(), nullptr, // readonly_block_num (test writes all) + nullptr, // finished RopeKernelParam{}, 2 * head_num * seq_len, 0, @@ -179,6 +180,7 @@ void TestBlocks(const thrust::universal_vector& k_cache, // [B, H, S, k_ptrs.data().get(), cu_seq_lens.data().get(), cu_block_cnts.data().get(), + nullptr, // finished RopeKernelParam{}, 2 * head_num * seq_len, 0, @@ -566,6 +568,7 @@ int test_attention() k_ptrs.data().get(), cu_kv_lens.data().get(), cu_block_cnts.data().get(), + nullptr, // finished RopeKernelParam{}, // DECODING ? nullptr : params.rope_theta, KvHeadNum * kContextLen, 0, diff --git a/src/turbomind/kernels/attention/verification/CMakeLists.txt b/src/turbomind/kernels/attention/verification/CMakeLists.txt new file mode 100644 index 0000000000..d41633453e --- /dev/null +++ b/src/turbomind/kernels/attention/verification/CMakeLists.txt @@ -0,0 +1,26 @@ +# Copyright (c) OpenMMLab. All rights reserved. + +set(TM_BUILD_VERIFICATION_ATTENTION_SM90 OFF) +if(NOT MSVC AND CMAKE_CUDA_COMPILER_VERSION VERSION_GREATER_EQUAL "12.0") + set(TM_BUILD_VERIFICATION_ATTENTION_SM90 ON) +endif() + +if(TM_BUILD_VERIFICATION_ATTENTION_SM90) + add_library(verification_attention STATIC dispatch.cu reduce.cu) + target_compile_definitions(verification_attention PUBLIC + TM_BUILD_VERIFICATION_ATTENTION_SM90=1) + set_property(TARGET verification_attention + PROPERTY CUDA_ARCHITECTURES 90a-real) + set_property(TARGET verification_attention + PROPERTY CUDA_RESOLVE_DEVICE_SYMBOLS ON) + target_compile_options(verification_attention PRIVATE + -O3 + $<$:-use_fast_math --expt-relaxed-constexpr -Xptxas=-v --threads 8>) + target_link_libraries(verification_attention PRIVATE nvidia::cutlass::cutlass) +else() + add_library(verification_attention STATIC stub.cc) + target_compile_definitions(verification_attention PUBLIC + TM_BUILD_VERIFICATION_ATTENTION_SM90=0) +endif() + +set_property(TARGET verification_attention PROPERTY POSITION_INDEPENDENT_CODE ON) diff --git a/src/turbomind/kernels/attention/verification/attention.h b/src/turbomind/kernels/attention/verification/attention.h new file mode 100644 index 0000000000..a6d9740171 --- /dev/null +++ b/src/turbomind/kernels/attention/verification/attention.h @@ -0,0 +1,125 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#pragma once + +#include + +#include + +#include + +#include "src/turbomind/core/data_type.h" +#include "src/turbomind/models/llama/llama_rope.h" + +namespace turbomind::verification_attention { + +struct Arguments { + void* out{}; + const void* q{}; + const void* q_bias{}; + int64_t q_stride{}; + + char* const* block_ptrs{}; + const int* block_ptr_offsets{}; + const int* q_offsets{}; + const int* k_offsets{}; + const bool* finished{}; + + int request_count{}; + int query_count{}; + int query_offset{}; + int max_query_length{}; + int max_key_length{}; + + int query_head_count{}; + int kv_head_count{}; + int query_group_size{}; + cutlass::FastDivmod query_group_size_divmod{}; + int head_dim{}; + int block_len{}; + cutlass::FastDivmod block_len_divmod{}; + int cache_block_offset{}; + int window_size{}; + + float qk_scale_log2{}; + RopeKernelParam rope{}; + + int split_count{}; + float* partial_o{}; + float* partial_ml{}; + + DataType data_type{}; + cudaStream_t stream{}; +}; + +struct Capability { + int arch{}; + DataType data_type{}; + int head_dim{}; + int max_query_length{}; + int quant_policy{}; + int cp_size{}; + bool is_mla{}; + bool has_attention_sinks{}; + int max_dynamic_smem_bytes{}; +}; + +bool supports(const Capability& capability); + +int choose_split_count(int query_count, + int base_cta_count, + int max_key_length, + int key_tile, + int partial_capacity, + int requested_max_splits, + int sm_count); + +inline bool UsesWgmma(const Arguments& arguments) +{ + return arguments.data_type == DataType::kBfloat16 + && arguments.head_dim == 256 + && arguments.block_len == 64 + && arguments.max_query_length <= 16 + && arguments.max_query_length * arguments.query_group_size <= 128; +} + +inline int WgmmaM(const Arguments& arguments) +{ + const int m = arguments.max_query_length * arguments.query_group_size; + return m <= 32 ? 32 : m <= 64 ? 64 : 128; +} + +inline bool UsesM32(const Arguments& arguments) +{ + return UsesWgmma(arguments) && WgmmaM(arguments) == 32; +} + +inline bool UsesM64(const Arguments& arguments) +{ + return UsesWgmma(arguments) && WgmmaM(arguments) == 64; +} + +inline bool UsesM128(const Arguments& arguments) +{ + return UsesWgmma(arguments) && WgmmaM(arguments) == 128; +} + +inline bool UsesRuntimeWgmmaM(const Arguments& arguments) +{ + return UsesM32(arguments) || UsesM64(arguments) || UsesM128(arguments); +} + +inline int CtaM(const Arguments& arguments) +{ + const int m = arguments.max_query_length * arguments.query_group_size; + return m <= 64 ? 64 : 128; +} + +inline int KeyTile(const Arguments& arguments) +{ + return arguments.head_dim == 128 || UsesRuntimeWgmmaM(arguments) ? 64 : 32; +} + +void run(const Arguments& arguments); + +} // namespace turbomind::verification_attention diff --git a/src/turbomind/kernels/attention/verification/dispatch.cu b/src/turbomind/kernels/attention/verification/dispatch.cu new file mode 100644 index 0000000000..86fef29e09 --- /dev/null +++ b/src/turbomind/kernels/attention/verification/dispatch.cu @@ -0,0 +1,110 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#include "src/turbomind/kernels/attention/verification/kernel_sm80.cuh" +#include "src/turbomind/kernels/attention/verification/kernel_sm90_wgmma.cuh" +#include "src/turbomind/kernels/attention/verification/kernel_sm90_wgmma_rs.cuh" +#include "src/turbomind/kernels/attention/verification/kernel_sm90_wgmma_rs_ws.cuh" + +#include + +namespace turbomind::verification_attention { + +bool supports(const Capability& c) +{ + const bool dtype_supported = + c.data_type == DataType::kHalf || c.data_type == DataType::kBfloat16; + const bool dimension_supported = c.head_dim == 128 || c.head_dim == 256; + const int required_dynamic_smem_bytes = c.head_dim == 128 ? 112 * 1024 : 104 * 1024; + return c.arch == 90 && dtype_supported && dimension_supported + && c.max_query_length <= 16 && c.quant_policy == 0 && c.cp_size == 1 + && !c.is_mla && !c.has_attention_sinks + && c.max_dynamic_smem_bytes >= required_dynamic_smem_bytes; +} + +int choose_split_count(int query_count, + int base_cta_count, + int max_key_length, + int key_tile, + int partial_capacity, + int requested_max_splits, + int sm_count) +{ + const int key_tiles = (max_key_length + key_tile - 1) / key_tile; + const int capacity_limit = std::max(1, partial_capacity / query_count); + const int useful_limit = std::min(key_tiles, requested_max_splits); + const int occupancy_target = std::max(1, sm_count * 2); + const int occupancy_splits = + (occupancy_target + base_cta_count - 1) / base_cta_count; + return std::min(128, + std::min(capacity_limit, + std::min(useful_limit, std::max(1, occupancy_splits)))); +} + +template +void Launch(const Arguments& arguments) +{ + const int m_count = arguments.max_query_length * arguments.query_group_size; + const int m_slices = (m_count + Policy::MTile - 1) / Policy::MTile; + const dim3 grid(arguments.request_count, + arguments.kv_head_count * m_slices, + arguments.split_count); + constexpr int smem_bytes = sizeof(SharedStorage); + if (arguments.split_count == 1) { + auto kernel = VerificationAttentionKernel; + cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_bytes); + cudaFuncSetAttribute(kernel, cudaFuncAttributePreferredSharedMemoryCarveout, 100); + kernel<<>>(arguments); + } + else { + auto kernel = VerificationAttentionKernel; + cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_bytes); + cudaFuncSetAttribute(kernel, cudaFuncAttributePreferredSharedMemoryCarveout, 100); + kernel<<>>(arguments); + } +} + +template +void DispatchM(const Arguments& arguments) +{ + if (CtaM(arguments) == 64) { + Launch>(arguments); + } + else { + Launch>(arguments); + } +} + +template +void DispatchHeadDimension(const Arguments& arguments) +{ + if (arguments.head_dim == 128) { + DispatchM(arguments); + } + else { + DispatchM(arguments); + } +} + +void run(const Arguments& arguments) +{ + if (UsesM128(arguments)) { + LaunchWgmmaRsWs(arguments); + } + else if (UsesM64(arguments)) { + LaunchWgmmaRs(arguments); + } + else if (UsesM32(arguments)) { + LaunchWgmma(arguments); + } + else if (arguments.data_type == DataType::kHalf) { + DispatchHeadDimension(arguments); + } + else { + DispatchHeadDimension(arguments); + } + if (arguments.split_count > 1) { + Reduce(arguments); + } +} + +} // namespace turbomind::verification_attention diff --git a/src/turbomind/kernels/attention/verification/kernel_sm80.cuh b/src/turbomind/kernels/attention/verification/kernel_sm80.cuh new file mode 100644 index 0000000000..381a9d3d17 --- /dev/null +++ b/src/turbomind/kernels/attention/verification/kernel_sm80.cuh @@ -0,0 +1,499 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#pragma once + +#include +#include +#include + +#include + +#include +#include + +#include "src/turbomind/kernels/attention/rotary_embedding.h" +#include "src/turbomind/kernels/attention/verification/attention.h" +#include "src/turbomind/kernels/attention/verification/paged_kv.cuh" +#include "src/turbomind/kernels/attention/verification/policy_sm90.cuh" +#include "src/turbomind/kernels/core/array_ops.h" + +namespace turbomind::verification_attention { + +// Generic mma.sync/ldmatrix fallback. It is compiled into the SM90a +// verification library but uses the SM80 tensor-core instruction family. + +CUTE_DEVICE float ReduceRowMax(float value) +{ + value = fmaxf(value, __shfl_xor_sync(0xffffffffu, value, 1)); + value = fmaxf(value, __shfl_xor_sync(0xffffffffu, value, 2)); + return value; +} + +CUTE_DEVICE float ReduceRowSum(float value) +{ + value += __shfl_xor_sync(0xffffffffu, value, 1); + value += __shfl_xor_sync(0xffffffffu, value, 2); + return value; +} + +template +struct VerificationAttentionMainloop { + Arguments arguments; + SharedStorage& storage; + + CUTE_DEVICE auto decode_m(int m, int m_begin) const + { + const int flat_m = m_begin + m; + int query_position; + int head_in_group; + arguments.query_group_size_divmod(query_position, head_in_group, flat_m); + return cute::make_coord(query_position, head_in_group); + } + + CUTE_DEVICE void run() + { + using Mma = Sm90Mma; + using TiledMmaQK = typename Mma::QK; + using TiledMmaPV = typename Mma::PV; + using QueryCopy = QTileCopy; + using KeyValueCopy = KvTileCopy; + + static_assert(KeyValueCopy::AccessCount == (Policy::Threads == 256 ? 4 : 8)); + static_assert(Policy::Stages == 3); + + using QLayout = SmemLayout2D; + using KVLayout = SmemLayout3D; + using PLayout = SmemLayout2D; + + auto shared_q = cute::make_tensor(cute::make_smem_ptr(storage.q), QLayout{}); + auto shared_k = cute::make_tensor(cute::make_smem_ptr(storage.body.k), KVLayout{}); + auto shared_v = cute::make_tensor(cute::make_smem_ptr(storage.body.v), KVLayout{}); + auto shared_probability = + cute::make_tensor(cute::make_smem_ptr(storage.body.probability), PLayout{}); + auto global_out = cute::make_tensor( + cute::make_gmem_ptr(static_cast(arguments.out)), + cute::make_layout( + cute::make_shape(arguments.query_count, + arguments.query_head_count, + cute::Int{}), + cute::make_stride(arguments.query_head_count * Policy::HeadDim, + cute::Int{}, + cute::_1{}))); + auto partial_o = cute::make_tensor( + cute::make_gmem_ptr(arguments.partial_o), + cute::make_layout( + cute::make_shape(arguments.query_count, + arguments.split_count, + arguments.query_head_count, + cute::Int{}), + cute::make_stride(arguments.split_count * arguments.query_head_count * Policy::HeadDim, + arguments.query_head_count * Policy::HeadDim, + cute::Int{}, + cute::_1{}))); + auto partial_ml = cute::make_tensor( + cute::make_gmem_ptr(arguments.partial_ml), + cute::make_layout( + cute::make_shape(arguments.query_count, + arguments.split_count, + arguments.query_head_count, + cute::_2{}), + cute::make_stride(arguments.split_count * arguments.query_head_count * 2, + arguments.query_head_count * 2, + cute::_2{}, + cute::_1{}))); + + const int request = blockIdx.x; + const int query_group_size = arguments.query_group_size; + const int m_count = arguments.max_query_length * query_group_size; + const int m_slices = (m_count + Policy::MTile - 1) / Policy::MTile; + auto block_head = cute::idx2crd( + static_cast(blockIdx.y), + cute::make_shape(m_slices, arguments.kv_head_count)); + const int m_begin = cute::get<0>(block_head) * Policy::MTile; + const int kv_head = cute::get<1>(block_head); + const int query_begin = arguments.q_offsets[request]; + const int query_end = arguments.q_offsets[request + 1]; + const int query_length = query_end - query_begin; + const int key_length = arguments.k_offsets[request + 1] - arguments.k_offsets[request]; + const int history_length = key_length - query_length; + const int tile_count = (key_length + Policy::KeyTile - 1) / Policy::KeyTile; + const int tiles_per_split = (tile_count + arguments.split_count - 1) / arguments.split_count; + const int first_tile = blockIdx.z * tiles_per_split; + const int last_tile = min(tile_count, first_tile + tiles_per_split); + const bool empty_split = first_tile >= last_tile; + const bool finished = arguments.finished && arguments.finished[request]; + + const bool inactive = finished || empty_split; + + auto thread_mma_pv = TiledMmaPV{}.get_slice(threadIdx.x); + auto pv_identity = cute::make_identity_tensor(typename Policy::PvShape{}); + auto pv_output_coordinates = thread_mma_pv.partition_C(pv_identity); + const int row0 = cute::get<0>(pv_output_coordinates(0)); + auto pv_output_prototype = thread_mma_pv.make_fragment_C(pv_output_coordinates); + using PvOutputLayout = typename decltype(pv_output_prototype)::layout_type; + static_assert(cute::rank_v == 3); + static constexpr int PvOutputValues = cute::cosize_v; + using PvTileMode = cute::Layout< + cute::Shape>, + cute::Stride>>; + using OutputFragmentsLayout = decltype(cute::append(PvOutputLayout{}, PvTileMode{})); + + cutlass::Array output_storage; + auto output_fragments = cute::make_tensor( + cute::make_rmem_ptr(output_storage.data()), OutputFragmentsLayout{}); + cute::clear(output_fragments); + + float running_max[2] = {-CUDART_INF_F, -CUDART_INF_F}; + float running_sum[2] = {0.f, 0.f}; + + if (!inactive) { + + auto global_q = cute::make_tensor( + cute::make_gmem_ptr(static_cast(arguments.q)), + cute::make_layout( + cute::make_shape(arguments.query_count, + arguments.query_head_count, + cute::Int{}), + cute::make_stride(arguments.q_stride, + cute::Int{}, + cute::_1{}))); + auto global_q_bias = cute::make_tensor( + cute::make_gmem_ptr(static_cast(arguments.q_bias)), + cute::make_layout( + cute::make_shape(arguments.query_head_count, + cute::Int{}), + cute::make_stride(cute::Int{}, cute::_1{}))); + + using RegisterCopy = cute::Copy_Atom, T>; + auto q_identity = cute::make_identity_tensor(typename Policy::QShape{}); + auto tiled_q_copy = typename QueryCopy::TiledCopy{}; + auto thread_q_copy = tiled_q_copy.get_thread_slice(threadIdx.x); + auto q_coordinates = cute::group_modes< + 1, cute::rank_v>( + thread_q_copy.partition_S(q_identity)); + auto q_destinations = cute::group_modes< + 1, cute::rank_v>( + thread_q_copy.partition_D(shared_q)); + Array fragment; + auto fragment_tensor = cute::make_tensor( + cute::make_rmem_ptr(fragment.data()), + cute::make_layout(cute::Int{})); + Array bias; + auto bias_fragment = cute::make_tensor( + cute::make_rmem_ptr(bias.data()), + cute::make_layout(cute::Int{})); + + CUTE_UNROLL + for (int access = 0; access < QueryCopy::AccessCount; ++access) { + const auto md = q_coordinates(cute::_0{}, access); + const auto qh = decode_m(cute::get<0>(md), m_begin); + const int query_position = cute::get<0>(qh); + const int head_in_group = cute::get<1>(qh); + const int d_begin = cute::get<1>(md); + const bool valid = query_position < query_length && head_in_group < query_group_size; + + CUTE_UNROLL + for (int value = 0; value < QueryCopy::ValuesPerAccess; ++value) { + fragment[value] = T(0); + } + if (valid) { + const int query_head = kv_head * query_group_size + head_in_group; + auto source = cute::make_tensor( + cute::make_gmem_ptr(&global_q(query_begin + query_position, query_head, d_begin)), + cute::make_layout(cute::Int{})); + cute::copy(RegisterCopy{}, source, fragment_tensor); + + if (arguments.q_bias) { + auto bias_source = cute::make_tensor( + cute::make_gmem_ptr(&global_q_bias(query_head, d_begin)), + cute::make_layout(cute::Int{})); + cute::copy(RegisterCopy{}, bias_source, bias_fragment); + CUTE_UNROLL + for (int value = 0; value < QueryCopy::ValuesPerAccess; ++value) { + fragment[value] = fragment[value] + bias[value]; + } + } + + FastRoPE rope( + arguments.rope, + request, + std::integral_constant{}); + rope.init(d_begin); + rope.apply(fragment, history_length + query_position); + + } + cute::copy(RegisterCopy{}, fragment_tensor, q_destinations(cute::_, access)); + } + __syncthreads(); + + auto thread_mma_qk = TiledMmaQK{}.get_slice(threadIdx.x); + auto tiled_copy_q = typename Mma::CopyQkA{}; + auto thread_copy_q = tiled_copy_q.get_slice(threadIdx.x); + auto tiled_copy_k = typename Mma::CopyQkB{}; + auto thread_copy_k = tiled_copy_k.get_slice(threadIdx.x); + auto tiled_copy_v = typename Mma::CopyPvB{}; + auto thread_copy_v = tiled_copy_v.get_slice(threadIdx.x); + + auto q_registers = thread_mma_qk.partition_fragment_A(shared_q); + auto q_source = thread_copy_q.partition_S(shared_q); + auto q_target = thread_copy_q.retile_D(q_registers); + cute::copy(tiled_copy_q, q_source, q_target); + // Q and K share storage; finish every Q read before staging K. + __syncthreads(); + + PagedKv cache(arguments.block_ptrs, + arguments.block_ptr_offsets, + request, + kv_head, + arguments.kv_head_count, + arguments.block_len, + arguments.block_len_divmod, + arguments.cache_block_offset); + + const int split_tile_count = last_tile - first_tile; + int issued_tiles = 0; + int consumed_tiles = 0; + int outstanding_groups = 0; + int read_stage = 0; + int write_stage = 0; + + CUTE_UNROLL + for (int prologue = 0; prologue < Policy::Stages - 1; ++prologue) { + if (prologue < split_tile_count) { + const int load_tile = first_tile + issued_tiles; + const int load_key_begin = load_tile * Policy::KeyTile; + auto shared_k_write = shared_k(cute::_, cute::_, write_stage); + auto shared_v_write = shared_v(cute::_, cute::_, write_stage); + copy_paged_tile(cache, load_key_begin, key_length, shared_k_write); + copy_paged_tile(cache, load_key_begin, key_length, shared_v_write); + cute::cp_async_fence(); + ++issued_tiles; + ++outstanding_groups; + write_stage = (write_stage + 1) % Policy::Stages; + } + } + + for (; consumed_tiles < split_tile_count; ++consumed_tiles) { + if (outstanding_groups == 1) { + cute::cp_async_wait<0>(); + } + else { + cute::cp_async_wait<1>(); + } + __syncthreads(); + --outstanding_groups; + + const int current_tile = first_tile + consumed_tiles; + const int current_key_begin = current_tile * Policy::KeyTile; + auto shared_k_read = shared_k(cute::_, cute::_, read_stage); + auto shared_v_read = shared_v(cute::_, cute::_, read_stage); + + if (issued_tiles < split_tile_count) { + const int load_tile = first_tile + issued_tiles; + const int load_key_begin = load_tile * Policy::KeyTile; + auto shared_k_write = shared_k(cute::_, cute::_, write_stage); + auto shared_v_write = shared_v(cute::_, cute::_, write_stage); + copy_paged_tile(cache, load_key_begin, key_length, shared_k_write); + copy_paged_tile(cache, load_key_begin, key_length, shared_v_write); + cute::cp_async_fence(); + ++issued_tiles; + ++outstanding_groups; + write_stage = (write_stage + 1) % Policy::Stages; + } + + auto score_identity = cute::make_identity_tensor(typename Policy::ScoreShape{}); + auto score_coordinates = thread_mma_qk.partition_C(score_identity); + auto score_fragment = thread_mma_qk.make_fragment_C(score_coordinates); + cute::clear(score_fragment); + + CUTE_UNROLL + for (int d_tile = 0; d_tile < Policy::HeadDim / 16; ++d_tile) { + auto key_tile = cute::local_tile( + shared_k_read, + cute::Shape, cute::_16>{}, + cute::make_coord(cute::_0{}, d_tile)); + auto key_fragment = thread_mma_qk.partition_fragment_B(key_tile); + cute::copy(tiled_copy_k, + thread_copy_k.partition_S(key_tile), + thread_copy_k.retile_D(key_fragment)); + auto q_slice = q_registers(cute::_, cute::_, d_tile); + auto q_single_k = cute::make_tensor(q_slice.data(), cute::append(q_slice.layout())); + cute::gemm(thread_mma_qk, q_single_k, key_fragment, score_fragment); + } + + float local_max[2] = {-CUDART_INF_F, -CUDART_INF_F}; + CUTE_UNROLL + for (int i = 0; i < cute::size(score_fragment); ++i) { + const auto mk = score_coordinates(i); + const auto qh = decode_m(cute::get<0>(mk), m_begin); + const int query_position = cute::get<0>(qh); + const int head_in_group = cute::get<1>(qh); + const int key_in_tile = cute::get<1>(mk); + const int absolute_key = current_key_begin + key_in_tile; + const int last_valid = history_length + query_position; + const int first_valid = max(0, last_valid - arguments.window_size + 1); + const bool valid = query_position < query_length && head_in_group < query_group_size + && absolute_key >= first_valid && absolute_key <= last_valid; + const float score = valid ? score_fragment(i) * arguments.qk_scale_log2 : + -CUDART_INF_F; + score_fragment(i) = score; + const int row_slot = cute::get<0>(mk) == row0 ? 0 : 1; + local_max[row_slot] = fmaxf(local_max[row_slot], score); + } + + float new_max[2]; + float old_scale[2]; + CUTE_UNROLL + for (int row = 0; row < 2; ++row) { + const float tile_max = ReduceRowMax(local_max[row]); + new_max[row] = fmaxf(running_max[row], tile_max); + old_scale[row] = running_max[row] == -CUDART_INF_F ? + 0.f : exp2f(running_max[row] - new_max[row]); + } + + CUTE_UNROLL + for (int pv_tile = 0; pv_tile < Policy::PvTileCount; ++pv_tile) { + auto output_tile = output_fragments(cute::_, cute::_, cute::_, pv_tile); + CUTE_UNROLL + for (int i = 0; i < cute::size(output_tile); ++i) { + const int row = cute::get<0>(pv_output_coordinates(i)); + output_tile(i) *= old_scale[row == row0 ? 0 : 1]; + } + } + + using ProbabilityStoreLayout = + typename decltype(cute::make_fragment_like(score_fragment))::layout_type; + auto probability_store_fragment = cute::make_tensor( + cute::recast_ptr(score_fragment.data()), ProbabilityStoreLayout{}); + float local_sum[2] = {0.f, 0.f}; + CUTE_UNROLL + for (int i = 0; i < cute::size(score_fragment); ++i) { + const int row = cute::get<0>(score_coordinates(i)); + const int row_slot = row == row0 ? 0 : 1; + const float score = score_fragment(i); + const float probability = score == -CUDART_INF_F ? + 0.f : exp2f(score - new_max[row_slot]); + probability_store_fragment(i) = static_cast(probability); + local_sum[row_slot] += probability; + } + + CUTE_UNROLL + for (int row = 0; row < 2; ++row) { + const float tile_sum = ReduceRowSum(local_sum[row]); + running_sum[row] = running_sum[row] * old_scale[row] + tile_sum; + running_max[row] = new_max[row]; + } + + auto store_p = typename Mma::StoreProbability{}; + auto thread_store_p = store_p.get_thread_slice(threadIdx.x); + cute::copy(store_p, + thread_store_p.retile_S(probability_store_fragment), + thread_store_p.partition_D(shared_probability)); + __syncthreads(); + + auto probability_fragment = thread_mma_pv.partition_fragment_A(shared_probability); + auto load_p = typename Mma::CopyPvA{}; + auto thread_load_p = load_p.get_thread_slice(threadIdx.x); + cute::copy(load_p, + thread_load_p.partition_S(shared_probability), + thread_load_p.retile_D(probability_fragment)); + + auto value_for_pv = cute::composition( + shared_v_read, + cute::Layout< + cute::Shape, cute::Int>, + cute::Stride, cute::_1>>{}); + CUTE_UNROLL + for (int pv_tile = 0; pv_tile < Policy::PvTileCount; ++pv_tile) { + auto output_tile = output_fragments(cute::_, cute::_, cute::_, pv_tile); + auto value_tile = cute::local_tile( + value_for_pv, + cute::Shape, cute::Int>{}, + cute::make_coord(pv_tile, cute::_0{})); + auto value_fragment = thread_mma_pv.partition_fragment_B(value_tile); + auto value_source = thread_copy_v.partition_S(value_tile); + auto value_target = thread_copy_v.retile_D(value_fragment); + CUTE_UNROLL + for (int k_block = 0; k_block < cute::size<2>(value_fragment); ++k_block) { + const auto k_coord = cute::idx2crd(k_block, cute::shape<2>(value_fragment)); + cute::copy(typename Mma::CopyPvBAtom{}, + value_source(cute::_, cute::_, k_coord), + value_target(cute::_, cute::_, k_coord)); + } + cute::gemm(thread_mma_pv, probability_fragment, value_fragment, output_tile); + } + + __syncthreads(); + read_stage = (read_stage + 1) % Policy::Stages; + } + } + + CUTE_UNROLL + for (int pv_tile = 0; pv_tile < Policy::PvTileCount; ++pv_tile) { + auto output_tile = output_fragments(cute::_, cute::_, cute::_, pv_tile); + CUTE_UNROLL + for (int i = 0; i < cute::size(output_tile); ++i) { + const auto mn = pv_output_coordinates(i); + const auto qh = decode_m(cute::get<0>(mn), m_begin); + const int query_position = cute::get<0>(qh); + const int head_in_group = cute::get<1>(qh); + if (query_position < query_length && head_in_group < query_group_size) { + const int absolute_query = query_begin + query_position; + const int query_head = kv_head * query_group_size + head_in_group; + const int d = cute::crd2idx( + cute::make_coord(cute::get<1>(mn), pv_tile), + cute::Shape, cute::Int>{}); + const int row_slot = cute::get<0>(mn) == row0 ? 0 : 1; + if constexpr (!StorePartial) { + global_out(absolute_query, query_head, d) = running_sum[row_slot] == 0.f ? + T(0) : static_cast(output_tile(i) / running_sum[row_slot]); + } + else { + const int local_query = absolute_query - arguments.query_offset; + partial_o(local_query, + static_cast(blockIdx.z), + query_head, + d) = output_tile(i); + } + } + } + } + + if constexpr (StorePartial) { + CUTE_UNROLL + for (int i = 0; i < cute::size(pv_output_coordinates); ++i) { + const auto mn = pv_output_coordinates(i); + if (cute::get<1>(mn) == 0) { + const auto qh = decode_m(cute::get<0>(mn), m_begin); + const int query_position = cute::get<0>(qh); + const int head_in_group = cute::get<1>(qh); + if (query_position < query_length && head_in_group < query_group_size) { + const int absolute_query = query_begin + query_position; + const int local_query = absolute_query - arguments.query_offset; + const int query_head = kv_head * query_group_size + head_in_group; + const int row_slot = cute::get<0>(mn) == row0 ? 0 : 1; + auto ml = partial_ml(local_query, + static_cast(blockIdx.z), + query_head, + cute::_); + ml(cute::_0{}) = running_max[row_slot]; + ml(cute::_1{}) = running_sum[row_slot]; + } + } + } + } + } +}; + +template +__global__ __maxnreg__(Policy::HeadDim == 128 ? 209 : 255) +void VerificationAttentionKernel(Arguments arguments) +{ + extern __shared__ char dynamic_shared[]; + auto& storage = *reinterpret_cast*>(dynamic_shared); + VerificationAttentionMainloop{arguments, storage}.run(); +} + +void Reduce(const Arguments& arguments); + +} // namespace turbomind::verification_attention diff --git a/src/turbomind/kernels/attention/verification/kernel_sm90_wgmma.cuh b/src/turbomind/kernels/attention/verification/kernel_sm90_wgmma.cuh new file mode 100644 index 0000000000..575ff42509 --- /dev/null +++ b/src/turbomind/kernels/attention/verification/kernel_sm90_wgmma.cuh @@ -0,0 +1,659 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#pragma once + +#include +#include + +#include + +#include +#include +#include +#include + +#include "src/turbomind/kernels/attention/rotary_embedding.h" +#include "src/turbomind/kernels/attention/verification/attention.h" +#include "src/turbomind/kernels/attention/verification/paged_kv.cuh" +#include "src/turbomind/kernels/core/array_ops.h" + +namespace turbomind::verification_attention { + +CUTE_DEVICE static void sync_named_barrier(int barrier_id, int thread_count) +{ + asm volatile("bar.sync %0, %1;" + : + : "r"(barrier_id), "r"(thread_count) + : "memory"); +} + +CUTE_DEVICE static void sync_warp_group_barrier(int warp_group) +{ + sync_named_barrier(8 + warp_group, 128); +} + +CUTE_DEVICE static void sync_compute_groups() +{ + sync_named_barrier(10, 256); +} + +struct Sm90WgmmaPolicy256 { + static constexpr int HeadDim = 256; + static constexpr int Threads = 128; + static constexpr int MTile = 32; + static constexpr int KeyTile = 64; + static constexpr int PvNtile = 64; + static constexpr int PvTileCount = HeadDim / PvNtile; + + using QShape = cute::Shape, cute::Int>; + using KvShape = cute::Shape, cute::Int>; +}; + +template +struct Sm90WgmmaMma; + +template<> +struct Sm90WgmmaMma { + using QK = decltype(cute::make_tiled_mma( + cute::SM90_64x32x16_F32BF16BF16_SS< + cute::GMMA::Major::K, cute::GMMA::Major::K>{})); + using PV = decltype(cute::make_tiled_mma( + cute::SM90_64x32x16_F32BF16BF16_SS< + cute::GMMA::Major::MN, cute::GMMA::Major::K>{})); +}; + +template +struct Sm90WgmmaKvStorage { + T k[Policy::KeyTile * Policy::HeadDim]; + T v[Policy::KeyTile * Policy::HeadDim]; +}; + +template +union alignas(128) Sm90WgmmaKvOrOutputStorage { + Sm90WgmmaKvStorage kv; + float output[Policy::MTile * Policy::HeadDim]; +}; + +template +struct Sm90WgmmaSplitStorage { + Sm90WgmmaKvOrOutputStorage kv_or_output; + alignas(128) T probability[Policy::MTile * Policy::KeyTile]; + alignas(16) float row_partials[4][Policy::MTile]; + alignas(16) float running_max[Policy::MTile]; + alignas(16) float running_sum[Policy::MTile]; + alignas(16) float old_scale[Policy::MTile]; +}; + +template +struct Sm90WgmmaSharedStorage { + alignas(128) T q[Policy::MTile * Policy::HeadDim]; + Sm90WgmmaSplitStorage split[2]; +}; + +template +struct Sm90WgmmaMainloop { + using Policy = Sm90WgmmaPolicy256; + + Arguments arguments; + Sm90WgmmaSharedStorage& storage; + + CUTE_DEVICE auto decode_row(int row) const + { + int query_position; + int head_in_group; + arguments.query_group_size_divmod(query_position, head_in_group, row); + return cute::make_coord(query_position, head_in_group); + } + + CUTE_DEVICE void run() + { + using Mma = Sm90WgmmaMma; + using QueryCopy = QTileCopy; + + using QLayout = decltype(cute::tile_to_shape( + cute::GMMA::Layout_K_SW128_Atom{}, + cute::Shape, cute::_256>{})); + using KLayout = decltype(cute::tile_to_shape( + cute::GMMA::Layout_K_SW128_Atom{}, + cute::Shape{})); + using VLayout = decltype(cute::tile_to_shape( + cute::GMMA::Layout_MN_SW128_Atom{}, + cute::Shape{})); + using PLayout = decltype(cute::tile_to_shape( + cute::GMMA::Layout_K_SW128_Atom{}, + cute::Shape, cute::_64>{})); + + const int warp_group = threadIdx.x / Policy::Threads; + const int local_tid = threadIdx.x % Policy::Threads; + const int split = StorePartial ? 2 * blockIdx.z + warp_group : warp_group; + const int work_split_count = StorePartial ? arguments.split_count : 2; + auto& split_storage = storage.split[warp_group]; + auto shared_q = cute::make_tensor(cute::make_smem_ptr(storage.q), QLayout{}); + auto shared_k = cute::make_tensor( + cute::make_smem_ptr(split_storage.kv_or_output.kv.k), KLayout{}); + auto shared_v = cute::make_tensor( + cute::make_smem_ptr(split_storage.kv_or_output.kv.v), VLayout{}); + auto shared_p = cute::make_tensor(cute::make_smem_ptr(split_storage.probability), PLayout{}); + auto v_copy_view = cute::composition( + shared_v, + cute::Layout< + cute::Shape, + cute::Stride, cute::_1>>{}); + + const int request = blockIdx.x; + const int query_group_size = arguments.query_group_size; + const int kv_head = blockIdx.y; + const int query_begin = arguments.q_offsets[request]; + const int query_end = arguments.q_offsets[request + 1]; + const int query_length = query_end - query_begin; + const int key_length = arguments.k_offsets[request + 1] - arguments.k_offsets[request]; + const int history_length = key_length - query_length; + const int page_count = (key_length + Policy::KeyTile - 1) / Policy::KeyTile; + const int first_page = static_cast(page_count) * split / work_split_count; + const int last_page = static_cast(page_count) * (split + 1) / work_split_count; + const bool request_finished = arguments.finished && arguments.finished[request]; + const bool inactive = request_finished || first_page >= last_page; + + auto global_out = cute::make_tensor( + cute::make_gmem_ptr(static_cast(arguments.out)), + cute::make_layout( + cute::make_shape(arguments.query_count, + arguments.query_head_count, + cute::_256{}), + cute::make_stride(arguments.query_head_count * Policy::HeadDim, + cute::Int{}, + cute::_1{}))); + auto partial_o = cute::make_tensor( + cute::make_gmem_ptr(arguments.partial_o), + cute::make_layout( + cute::make_shape(arguments.query_count, + arguments.split_count, + arguments.query_head_count, + cute::_256{}), + cute::make_stride(arguments.split_count * arguments.query_head_count * Policy::HeadDim, + arguments.query_head_count * Policy::HeadDim, + cute::Int{}, + cute::_1{}))); + auto partial_ml = cute::make_tensor( + cute::make_gmem_ptr(arguments.partial_ml), + cute::make_layout( + cute::make_shape(arguments.query_count, + arguments.split_count, + arguments.query_head_count, + cute::_2{}), + cute::make_stride(arguments.split_count * arguments.query_head_count * 2, + arguments.query_head_count * 2, + cute::_2{}, + cute::_1{}))); + + auto pv_mma = typename Mma::PV{}; + auto thread_pv = pv_mma.get_slice(local_tid); + auto pv_identity = cute::make_identity_tensor( + cute::Shape>{}); + auto pv_coordinates = thread_pv.partition_C(pv_identity); + auto pv_prototype = thread_pv.make_fragment_C(pv_coordinates); + using PvLayout = typename decltype(pv_prototype)::layout_type; + static constexpr int PvValues = cute::cosize_v; + using PvTiles = cute::Layout< + cute::Shape>, + cute::Stride>>; + using OutputLayout = decltype(cute::append(PvLayout{}, PvTiles{})); + cutlass::Array output_storage; + auto output = cute::make_tensor( + cute::make_rmem_ptr(output_storage.data()), OutputLayout{}); + cute::clear(output); + + if (local_tid < Policy::MTile) { + split_storage.running_max[local_tid] = -CUDART_INF_F; + split_storage.running_sum[local_tid] = 0.f; + split_storage.old_scale[local_tid] = 0.f; + } + sync_warp_group_barrier(warp_group); + + if (warp_group == 0 && !request_finished) { + auto global_q = cute::make_tensor( + cute::make_gmem_ptr(static_cast(arguments.q)), + cute::make_layout( + cute::make_shape(arguments.query_count, + arguments.query_head_count, + cute::_256{}), + cute::make_stride(arguments.q_stride, + cute::Int{}, + cute::_1{}))); + auto global_q_bias = cute::make_tensor( + cute::make_gmem_ptr(static_cast(arguments.q_bias)), + cute::make_layout( + cute::make_shape(arguments.query_head_count, cute::_256{}), + cute::make_stride(cute::Int{}, cute::_1{}))); + + using RegisterCopy = cute::Copy_Atom, T>; + auto q_identity = cute::make_identity_tensor(typename Policy::QShape{}); + auto tiled_q_copy = typename QueryCopy::TiledCopy{}; + auto thread_q_copy = tiled_q_copy.get_thread_slice(local_tid); + auto q_coordinates = cute::group_modes< + 1, cute::rank_v>( + thread_q_copy.partition_S(q_identity)); + auto q_destinations = cute::group_modes< + 1, cute::rank_v>( + thread_q_copy.partition_D(shared_q)); + Array fragment; + auto fragment_tensor = cute::make_tensor( + cute::make_rmem_ptr(fragment.data()), + cute::make_layout(cute::Int{})); + Array bias; + auto bias_fragment = cute::make_tensor( + cute::make_rmem_ptr(bias.data()), + cute::make_layout(cute::Int{})); + + CUTE_UNROLL + for (int access = 0; access < QueryCopy::AccessCount; ++access) { + const auto nd = q_coordinates(cute::_0{}, access); + const auto qh = decode_row(cute::get<0>(nd)); + const int query_position = cute::get<0>(qh); + const int head_in_group = cute::get<1>(qh); + const int d_begin = cute::get<1>(nd); + const bool valid = query_position < query_length + && head_in_group < query_group_size; + CUTE_UNROLL + for (int value = 0; value < QueryCopy::ValuesPerAccess; ++value) { + fragment[value] = T(0); + } + if (valid) { + const int query_head = kv_head * query_group_size + head_in_group; + auto source = cute::make_tensor( + cute::make_gmem_ptr(&global_q(query_begin + query_position, + query_head, + d_begin)), + cute::make_layout(cute::Int{})); + cute::copy(RegisterCopy{}, source, fragment_tensor); + if (arguments.q_bias) { + auto bias_source = cute::make_tensor( + cute::make_gmem_ptr(&global_q_bias(query_head, d_begin)), + cute::make_layout(cute::Int{})); + cute::copy(RegisterCopy{}, bias_source, bias_fragment); + CUTE_UNROLL + for (int value = 0; value < QueryCopy::ValuesPerAccess; ++value) { + fragment[value] = fragment[value] + bias[value]; + } + } + FastRoPE rope( + arguments.rope, + request, + std::integral_constant{}); + rope.init(d_begin); + rope.apply(fragment, history_length + query_position); + } + cute::copy(RegisterCopy{}, fragment_tensor, q_destinations(cute::_, access)); + } + cutlass::arch::fence_view_async_shared(); + } + __syncthreads(); + + if constexpr (StorePartial) { + if (split >= arguments.split_count) { + return; + } + } + + if (!inactive) { + + PagedKv cache(arguments.block_ptrs, + arguments.block_ptr_offsets, + request, + kv_head, + arguments.kv_head_count, + arguments.block_len, + arguments.block_len_divmod, + arguments.cache_block_offset); + + auto qk_mma = typename Mma::QK{}; + auto thread_qk = qk_mma.get_slice(local_tid); + auto qk_identity = cute::make_identity_tensor( + cute::Shape>{}); + auto score_coordinates = thread_qk.partition_C(qk_identity); + + copy_paged_page( + cache, first_page, first_page * Policy::KeyTile, key_length, shared_k, local_tid); + cute::cp_async_fence(); + copy_paged_page( + cache, first_page, first_page * Policy::KeyTile, key_length, v_copy_view, local_tid); + cute::cp_async_fence(); + + for (int page = first_page; page < last_page; ++page) { + const int key_begin = page * Policy::KeyTile; + cute::cp_async_wait<1>(); + sync_warp_group_barrier(warp_group); + + auto score = thread_qk.make_fragment_C(score_coordinates); + cute::clear(score); + auto k_source = thread_qk.partition_A(shared_k); + auto q_source = thread_qk.partition_B(shared_q); + auto k_fragment = thread_qk.make_fragment_A(k_source); + auto q_fragment = thread_qk.make_fragment_B(q_source); + cute::warpgroup_fence_operand(score); + cute::warpgroup_arrive(); + cute::gemm(qk_mma, k_fragment, q_fragment, score); + cute::warpgroup_commit_batch(); + cute::warpgroup_wait<0>(); + cute::warpgroup_fence_operand(score); + + const bool has_next_page = page + 1 < last_page; + if (has_next_page) { + const int next_page = page + 1; + copy_paged_page( + cache, + next_page, + next_page * Policy::KeyTile, + key_length, + shared_k, + local_tid); + cute::cp_async_fence(); + } + + static constexpr int RowSlots = Policy::MTile / 4; + Array row_max; + CUTE_UNROLL + for (int row_slot = 0; row_slot < RowSlots; ++row_slot) { + row_max[row_slot] = -CUDART_INF_F; + } + CUTE_UNROLL + for (int i = 0; i < cute::size(score); ++i) { + const auto kn = score_coordinates(i); + const int key_in_page = cute::get<0>(kn); + const int row = cute::get<1>(kn); + const auto qh = decode_row(row); + const int query_position = cute::get<0>(qh); + const int head_in_group = cute::get<1>(qh); + const int absolute_key = key_begin + key_in_page; + const int last_valid = history_length + query_position; + const int first_valid = max(0, last_valid - arguments.window_size + 1); + const bool valid = query_position < query_length + && head_in_group < query_group_size + && absolute_key >= first_valid + && absolute_key <= last_valid; + score(i) = valid ? score(i) * arguments.qk_scale_log2 : + -CUDART_INF_F; + const int row_slot = (i / 4) * 2 + i % 2; + row_max[row_slot] = fmaxf(row_max[row_slot], score(i)); + } + + CUTE_UNROLL + for (int offset = 4; offset <= 16; offset *= 2) { + CUTE_UNROLL + for (int row_slot = 0; row_slot < RowSlots; ++row_slot) { + row_max[row_slot] = fmaxf( + row_max[row_slot], + __shfl_xor_sync(0xffffffffu, + row_max[row_slot], + offset)); + } + } + const int lane = local_tid % 32; + const int warp = local_tid / 32; + if (lane < 4) { + CUTE_UNROLL + for (int row_slot = 0; row_slot < RowSlots; ++row_slot) { + const int fragment_index = (row_slot / 2) * 4 + row_slot % 2; + const int row = cute::get<1>(score_coordinates(fragment_index)); + split_storage.row_partials[warp][row] = row_max[row_slot]; + } + } + sync_warp_group_barrier(warp_group); + + if (local_tid < Policy::MTile) { + const int row = local_tid; + float page_max = -CUDART_INF_F; + CUTE_UNROLL + for (int source_warp = 0; source_warp < 4; ++source_warp) { + page_max = fmaxf(page_max, + split_storage.row_partials[source_warp][row]); + } + const float old_max = split_storage.running_max[row]; + const float new_max = fmaxf(old_max, page_max); + const float scale = old_max == -CUDART_INF_F ? + 0.f : exp2f(old_max - new_max); + split_storage.old_scale[row] = scale; + split_storage.running_max[row] = new_max; + } + sync_warp_group_barrier(warp_group); + + Array row_sum; + CUTE_UNROLL + for (int row_slot = 0; row_slot < RowSlots; ++row_slot) { + row_sum[row_slot] = 0.f; + } + CUTE_UNROLL + for (int i = 0; i < cute::size(score); ++i) { + const int row = cute::get<1>(score_coordinates(i)); + const float value = score(i); + const float probability = value == -CUDART_INF_F ? + 0.f : + exp2f(value - split_storage.running_max[row]); + score(i) = probability; + const int row_slot = (i / 4) * 2 + i % 2; + row_sum[row_slot] += probability; + } + CUTE_UNROLL + for (int offset = 4; offset <= 16; offset *= 2) { + CUTE_UNROLL + for (int row_slot = 0; row_slot < RowSlots; ++row_slot) { + row_sum[row_slot] += __shfl_xor_sync( + 0xffffffffu, row_sum[row_slot], offset); + } + } + if (lane < 4) { + CUTE_UNROLL + for (int row_slot = 0; row_slot < RowSlots; ++row_slot) { + const int fragment_index = (row_slot / 2) * 4 + row_slot % 2; + const int row = cute::get<1>(score_coordinates(fragment_index)); + split_storage.row_partials[warp][row] = row_sum[row_slot]; + } + } + sync_warp_group_barrier(warp_group); + + if (local_tid < Policy::MTile) { + const int row = local_tid; + float page_sum = 0.f; + CUTE_UNROLL + for (int source_warp = 0; source_warp < 4; ++source_warp) { + page_sum += split_storage.row_partials[source_warp][row]; + } + split_storage.running_sum[row] = + split_storage.running_sum[row] * split_storage.old_scale[row] + page_sum; + } + sync_warp_group_barrier(warp_group); + + CUTE_UNROLL + for (int i = 0; i < cute::size(score); ++i) { + const auto kn = score_coordinates(i); + shared_p(cute::get<1>(kn), cute::get<0>(kn)) = + static_cast(score(i)); + } + cutlass::arch::fence_view_async_shared(); + if (has_next_page) { + cute::cp_async_wait<1>(); + } + else { + cute::cp_async_wait<0>(); + } + sync_warp_group_barrier(warp_group); + + CUTE_UNROLL + for (int tile = 0; tile < Policy::PvTileCount; ++tile) { + auto output_tile = output(cute::_, cute::_, cute::_, tile); + CUTE_UNROLL + for (int i = 0; i < cute::size(output_tile); ++i) { + const int row = cute::get<1>(pv_coordinates(i)); + output_tile(i) *= split_storage.old_scale[row]; + } + cute::warpgroup_fence_operand(output_tile); + } + cute::warpgroup_arrive(); + CUTE_UNROLL + for (int tile = 0; tile < Policy::PvTileCount; ++tile) { + auto output_tile = output(cute::_, cute::_, cute::_, tile); + auto value_tile = cute::local_tile( + shared_v, + cute::Shape{}, + cute::make_coord(tile, cute::_0{})); + auto value_source = thread_pv.partition_A(value_tile); + auto probability_source = thread_pv.partition_B(shared_p); + auto value_fragment = thread_pv.make_fragment_A(value_source); + auto probability_fragment = thread_pv.make_fragment_B(probability_source); + cute::gemm(pv_mma, + value_fragment, + probability_fragment, + output_tile); + } + cute::warpgroup_commit_batch(); + cute::warpgroup_wait<0>(); + CUTE_UNROLL + for (int tile = 0; tile < Policy::PvTileCount; ++tile) { + auto output_tile = output(cute::_, cute::_, cute::_, tile); + cute::warpgroup_fence_operand(output_tile); + } + if (has_next_page) { + const int next_page = page + 1; + copy_paged_page( + cache, + next_page, + next_page * Policy::KeyTile, + key_length, + v_copy_view, + local_tid); + cute::cp_async_fence(); + } + } + } + + CUTE_UNROLL + for (int tile = 0; tile < Policy::PvTileCount; ++tile) { + auto output_tile = output(cute::_, cute::_, cute::_, tile); + CUTE_UNROLL + for (int i = 0; i < cute::size(output_tile); ++i) { + const auto dn = pv_coordinates(i); + const int row = cute::get<1>(dn); + const auto qh = decode_row(row); + const int query_position = cute::get<0>(qh); + const int head_in_group = cute::get<1>(qh); + if (query_position < query_length && head_in_group < query_group_size) { + const int absolute_query = query_begin + query_position; + const int query_head = kv_head * query_group_size + head_in_group; + const int d = tile * Policy::PvNtile + cute::get<0>(dn); + if constexpr (StorePartial) { + const int local_query = absolute_query - arguments.query_offset; + partial_o(local_query, split, query_head, d) = output_tile(i); + } + else { + auto merge_output = cute::make_tensor( + cute::make_smem_ptr(split_storage.kv_or_output.output), + cute::make_layout(cute::Shape, cute::_256>{}, + cute::Stride{})); + merge_output(row, d) = output_tile(i); + } + } + } + } + + if constexpr (StorePartial) { + if (local_tid < Policy::MTile) { + const int row = local_tid; + const auto qh = decode_row(row); + const int query_position = cute::get<0>(qh); + const int head_in_group = cute::get<1>(qh); + if (query_position < query_length && head_in_group < query_group_size) { + const int absolute_query = query_begin + query_position; + const int local_query = absolute_query - arguments.query_offset; + const int query_head = kv_head * query_group_size + head_in_group; + partial_ml(local_query, split, query_head, cute::_0{}) = + split_storage.running_max[row]; + partial_ml(local_query, split, query_head, cute::_1{}) = + split_storage.running_sum[row]; + } + } + } + else { + __syncthreads(); + if (warp_group != 0) { + return; + } + + if (local_tid < Policy::MTile) { + const int row = local_tid; + const float max0 = storage.split[0].running_max[row]; + const float max1 = storage.split[1].running_max[row]; + const float merged_max = fmaxf(max0, max1); + const float scale0 = max0 == -CUDART_INF_F ? 0.f : exp2f(max0 - merged_max); + const float scale1 = max1 == -CUDART_INF_F ? 0.f : exp2f(max1 - merged_max); + const float merged_sum = storage.split[0].running_sum[row] * scale0 + + storage.split[1].running_sum[row] * scale1; + storage.split[0].old_scale[row] = scale0; + storage.split[1].old_scale[row] = scale1; + storage.split[0].running_sum[row] = merged_sum == 0.f ? 0.f : 1.f / merged_sum; + } + sync_warp_group_barrier(warp_group); + + auto output0 = cute::make_tensor( + cute::make_smem_ptr(storage.split[0].kv_or_output.output), + cute::make_layout(cute::Shape, cute::_256>{}, + cute::Stride{})); + auto output1 = cute::make_tensor( + cute::make_smem_ptr(storage.split[1].kv_or_output.output), + cute::make_layout(cute::Shape, cute::_256>{}, + cute::Stride{})); + CUTE_UNROLL + for (int access = 0; + access < Policy::MTile * Policy::HeadDim / Policy::Threads; + ++access) { + const int linear = local_tid + access * Policy::Threads; + const int row = linear / Policy::HeadDim; + const int d = linear % Policy::HeadDim; + const auto qh = decode_row(row); + const int query_position = cute::get<0>(qh); + const int head_in_group = cute::get<1>(qh); + if (query_position < query_length && head_in_group < query_group_size) { + const float value = output0(row, d) * storage.split[0].old_scale[row] + + output1(row, d) * storage.split[1].old_scale[row]; + const int absolute_query = query_begin + query_position; + const int query_head = kv_head * query_group_size + head_in_group; + global_out(absolute_query, query_head, d) = static_cast( + value * storage.split[0].running_sum[row]); + } + } + } + } +}; + +template +__global__ __launch_bounds__(256, 1) +void VerificationAttentionWgmmaKernel(Arguments arguments) +{ + using Policy = Sm90WgmmaPolicy256; + extern __shared__ char dynamic_shared[]; + auto& storage = *reinterpret_cast*>(dynamic_shared); + Sm90WgmmaMainloop{arguments, storage}.run(); +} + +template +void LaunchWgmma(const Arguments& arguments) +{ + using Policy = Sm90WgmmaPolicy256; + const dim3 grid(arguments.request_count, + arguments.kv_head_count, + (arguments.split_count + 1) / 2); + constexpr int smem_bytes = sizeof(Sm90WgmmaSharedStorage); + if (arguments.split_count == 1) { + auto kernel = VerificationAttentionWgmmaKernel; + cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_bytes); + cudaFuncSetAttribute(kernel, cudaFuncAttributePreferredSharedMemoryCarveout, 100); + kernel<<>>(arguments); + } + else { + auto kernel = VerificationAttentionWgmmaKernel; + cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_bytes); + cudaFuncSetAttribute(kernel, cudaFuncAttributePreferredSharedMemoryCarveout, 100); + kernel<<>>(arguments); + } +} + +} // namespace turbomind::verification_attention diff --git a/src/turbomind/kernels/attention/verification/kernel_sm90_wgmma_rs.cuh b/src/turbomind/kernels/attention/verification/kernel_sm90_wgmma_rs.cuh new file mode 100644 index 0000000000..6541466797 --- /dev/null +++ b/src/turbomind/kernels/attention/verification/kernel_sm90_wgmma_rs.cuh @@ -0,0 +1,776 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#pragma once + +#include "src/turbomind/kernels/attention/verification/kernel_sm90_wgmma.cuh" + +namespace turbomind::verification_attention { + +struct Sm90WgmmaRsPolicy256 { + static constexpr int HeadDim = 256; + static constexpr int Threads = 128; + static constexpr int MTile = 64; + static constexpr int KeyTile = 64; + static constexpr int QueryPartitionTile = 8; + + using QShape = cute::Shape, cute::Int>; + using KvShape = cute::Shape, cute::Int>; +}; + +template +struct Sm90WgmmaRsMma; + +template<> +struct Sm90WgmmaRsMma { + using QK = decltype(cute::make_tiled_mma( + cute::SM90_64x64x16_F32BF16BF16_SS< + cute::GMMA::Major::K, cute::GMMA::Major::K>{})); + using PV = decltype(cute::make_tiled_mma( + cute::SM90_64x256x16_F32BF16BF16_RS< + cute::GMMA::Major::K, cute::GMMA::Major::MN>{})); +}; + +template +struct Sm90WgmmaRsKvStorage { + using Policy = Sm90WgmmaRsPolicy256; + T k[Policy::KeyTile * Policy::HeadDim]; + T v[Policy::KeyTile * Policy::HeadDim]; +}; + +template +union alignas(128) Sm90WgmmaRsGroupBody { + using Policy = Sm90WgmmaRsPolicy256; + Sm90WgmmaRsKvStorage kv; + float output[Policy::MTile * Policy::HeadDim]; +}; + +template +struct Sm90WgmmaRsGroupStorage { + using Policy = Sm90WgmmaRsPolicy256; + Sm90WgmmaRsGroupBody body; + alignas(16) float running_max[Policy::MTile]; + alignas(16) float running_sum[Policy::MTile]; +}; + +template +struct Sm90WgmmaRsSharedStorage { + using Policy = Sm90WgmmaRsPolicy256; + alignas(128) T q[2][Policy::MTile * Policy::HeadDim]; + Sm90WgmmaRsGroupStorage group[2]; +}; + +template +struct Sm90WgmmaRsMainloop { + using Policy = Sm90WgmmaRsPolicy256; + using Mma = Sm90WgmmaRsMma; + + Arguments arguments; + Sm90WgmmaRsSharedStorage& storage; + + CUTE_DEVICE auto decode_row(int row) const + { + int query_position; + int head_in_group; + arguments.query_group_size_divmod(query_position, head_in_group, row); + return cute::make_coord(query_position, head_in_group); + } + + CUTE_DEVICE void run() + { + // The registered M64 path uses two warp groups over disjoint context + // ranges and merges their online-softmax states. + using QueryCopy = QTileCopy; + using QLayout = decltype(cute::tile_to_shape( + cute::GMMA::Layout_K_SW128_Atom{}, + cute::Shape{})); + using KLayout = QLayout; + using VLayout = decltype(cute::tile_to_shape( + cute::GMMA::Layout_MN_SW128_Atom{}, + cute::Shape{})); + + const int warp_group = threadIdx.x / Policy::Threads; + const int local_tid = threadIdx.x % Policy::Threads; + const int query_base = QueryPartition ? warp_group * Policy::QueryPartitionTile : 0; + const int split = QueryPartition ? blockIdx.z : + (StorePartial ? 2 * blockIdx.z + warp_group : warp_group); + const int work_split_count = QueryPartition ? arguments.split_count : + (StorePartial ? arguments.split_count : 2); + auto& group_storage = storage.group[warp_group]; + auto shared_q = cute::make_tensor( + cute::make_smem_ptr(storage.q[QueryPartition ? warp_group : 0]), QLayout{}); + auto& kv_storage = QueryPartition ? storage.group[0] : group_storage; + auto shared_k = cute::make_tensor(cute::make_smem_ptr(kv_storage.body.kv.k), KLayout{}); + auto shared_v = cute::make_tensor(cute::make_smem_ptr(kv_storage.body.kv.v), VLayout{}); + auto v_copy_view = cute::composition( + shared_v, + cute::Layout< + cute::Shape, + cute::Stride, cute::_1>>{}); + + const int request = blockIdx.x; + const int query_group_size = arguments.query_group_size; + const int kv_head = blockIdx.y; + const int query_begin = arguments.q_offsets[request]; + const int query_end = arguments.q_offsets[request + 1]; + const int query_length = query_end - query_begin; + const int key_length = arguments.k_offsets[request + 1] - arguments.k_offsets[request]; + const int history_length = key_length - query_length; + const int page_count = (key_length + Policy::KeyTile - 1) / Policy::KeyTile; + const int first_page = static_cast(page_count) * split / work_split_count; + const int last_page = static_cast(page_count) * (split + 1) / work_split_count; + const bool request_finished = arguments.finished && arguments.finished[request]; + const bool inactive = request_finished || first_page >= last_page; + + auto global_out = cute::make_tensor( + cute::make_gmem_ptr(static_cast(arguments.out)), + cute::make_layout( + cute::make_shape(arguments.query_count, + arguments.query_head_count, + cute::_256{}), + cute::make_stride(arguments.query_head_count * Policy::HeadDim, + cute::Int{}, + cute::_1{}))); + auto partial_o = cute::make_tensor( + cute::make_gmem_ptr(arguments.partial_o), + cute::make_layout( + cute::make_shape(arguments.query_count, + arguments.split_count, + arguments.query_head_count, + cute::_256{}), + cute::make_stride(arguments.split_count * arguments.query_head_count * Policy::HeadDim, + arguments.query_head_count * Policy::HeadDim, + cute::Int{}, + cute::_1{}))); + auto partial_ml = cute::make_tensor( + cute::make_gmem_ptr(arguments.partial_ml), + cute::make_layout( + cute::make_shape(arguments.query_count, + arguments.split_count, + arguments.query_head_count, + cute::_2{}), + cute::make_stride(arguments.split_count * arguments.query_head_count * 2, + arguments.query_head_count * 2, + cute::_2{}, + cute::_1{}))); + + auto pv_mma = typename Mma::PV{}; + auto thread_pv = pv_mma.get_slice(local_tid); + auto pv_identity = cute::make_identity_tensor(cute::Shape{}); + auto pv_coordinates = thread_pv.partition_C(pv_identity); + const int row0 = cute::get<0>(pv_coordinates(0)); + const int row1 = cute::get<0>(pv_coordinates(2)); + auto output = thread_pv.make_fragment_C(pv_coordinates); + cute::clear(output); + using ProbabilityLayout = cute::Layout< + cute::Shape, cute::_1, cute::_4>, + cute::Stride, cute::_0, cute::_8>>; + cutlass::Array probability_storage; + auto probability = cute::make_tensor( + cute::make_rmem_ptr(probability_storage.data()), ProbabilityLayout{}); + + Array running_max; + Array running_sum; + Array old_scale; + CUTE_UNROLL + for (int row_slot = 0; row_slot < 2; ++row_slot) { + running_max[row_slot] = -CUDART_INF_F; + running_sum[row_slot] = 0.f; + old_scale[row_slot] = 0.f; + } + + if ((QueryPartition || warp_group == 0) && !request_finished) { + auto global_q = cute::make_tensor( + cute::make_gmem_ptr(static_cast(arguments.q)), + cute::make_layout( + cute::make_shape(arguments.query_count, + arguments.query_head_count, + cute::_256{}), + cute::make_stride(arguments.q_stride, + cute::Int{}, + cute::_1{}))); + auto global_q_bias = cute::make_tensor( + cute::make_gmem_ptr(static_cast(arguments.q_bias)), + cute::make_layout( + cute::make_shape(arguments.query_head_count, cute::_256{}), + cute::make_stride(cute::Int{}, cute::_1{}))); + using RegisterCopy = cute::Copy_Atom, T>; + auto q_identity = cute::make_identity_tensor(typename Policy::QShape{}); + auto tiled_q_copy = typename QueryCopy::TiledCopy{}; + auto thread_q_copy = tiled_q_copy.get_thread_slice(local_tid); + auto q_coordinates = cute::group_modes< + 1, cute::rank_v>( + thread_q_copy.partition_S(q_identity)); + auto q_destinations = cute::group_modes< + 1, cute::rank_v>( + thread_q_copy.partition_D(shared_q)); + Array fragment; + auto fragment_tensor = cute::make_tensor( + cute::make_rmem_ptr(fragment.data()), + cute::make_layout(cute::Int{})); + Array bias; + auto bias_fragment = cute::make_tensor( + cute::make_rmem_ptr(bias.data()), + cute::make_layout(cute::Int{})); + + CUTE_UNROLL + for (int access = 0; access < QueryCopy::AccessCount; ++access) { + const auto nd = q_coordinates(cute::_0{}, access); + const auto qh = decode_row(cute::get<0>(nd)); + const int query_position = query_base + cute::get<0>(qh); + const int head_in_group = cute::get<1>(qh); + const int d_begin = cute::get<1>(nd); + const bool valid = query_position < query_length + && head_in_group < query_group_size; + CUTE_UNROLL + for (int value = 0; value < QueryCopy::ValuesPerAccess; ++value) { + fragment[value] = T(0); + } + if (valid) { + const int query_head = kv_head * query_group_size + head_in_group; + auto source = cute::make_tensor( + cute::make_gmem_ptr(&global_q(query_begin + query_position, + query_head, + d_begin)), + cute::make_layout(cute::Int{})); + cute::copy(RegisterCopy{}, source, fragment_tensor); + if (arguments.q_bias) { + auto bias_source = cute::make_tensor( + cute::make_gmem_ptr(&global_q_bias(query_head, d_begin)), + cute::make_layout(cute::Int{})); + cute::copy(RegisterCopy{}, bias_source, bias_fragment); + CUTE_UNROLL + for (int value = 0; value < QueryCopy::ValuesPerAccess; ++value) { + fragment[value] = fragment[value] + bias[value]; + } + } + FastRoPE rope( + arguments.rope, + request, + std::integral_constant{}); + rope.init(d_begin); + rope.apply(fragment, history_length + query_position); + } + cute::copy(RegisterCopy{}, fragment_tensor, q_destinations(cute::_, access)); + } + cutlass::arch::fence_view_async_shared(); + } + __syncthreads(); + + if constexpr (StorePartial) { + if (split >= arguments.split_count) { + return; + } + } + + if constexpr (QueryPartition) { + if (!inactive) { + PagedKv cache(arguments.block_ptrs, + arguments.block_ptr_offsets, + request, + kv_head, + arguments.kv_head_count, + arguments.block_len, + arguments.block_len_divmod, + arguments.cache_block_offset); + auto qk_mma = typename Mma::QK{}; + auto thread_qk = qk_mma.get_slice(local_tid); + auto qk_identity = cute::make_identity_tensor(cute::Shape{}); + auto score_coordinates = thread_qk.partition_C(qk_identity); + + const int initial_page = first_page + warp_group; + if (initial_page < last_page) { + auto initial_k = cute::make_tensor( + cute::make_smem_ptr(storage.group[warp_group].body.kv.k), KLayout{}); + auto initial_v = cute::make_tensor( + cute::make_smem_ptr(storage.group[warp_group].body.kv.v), VLayout{}); + auto initial_v_copy = cute::composition( + initial_v, + cute::Layout< + cute::Shape, + cute::Stride, cute::_1>>{}); + copy_paged_page( + cache, + initial_page, + initial_page * Policy::KeyTile, + key_length, + initial_k, + local_tid); + cute::cp_async_fence(); + copy_paged_page( + cache, + initial_page, + initial_page * Policy::KeyTile, + key_length, + initial_v_copy, + local_tid); + cute::cp_async_fence(); + } + + for (int pair_page = first_page; pair_page < last_page; pair_page += 2) { + // K-ready, V-ready, K-reuse, and V-reuse are the four CTA + // rendezvous for a nonfinal pair. The final two omit reuse. + const bool has_second_page = pair_page + 1 < last_page; + const bool has_next_pair = pair_page + 2 < last_page; + const int next_page = pair_page + 2 + warp_group; + + cute::cp_async_wait<1>(); + __syncthreads(); + + CUTE_UNROLL + for (int slot = 0; slot < 2; ++slot) { + if (slot == 0 || has_second_page) { + const int page = pair_page + slot; + const int key_begin = page * Policy::KeyTile; + auto shared_k_slot = cute::make_tensor( + cute::make_smem_ptr(storage.group[slot].body.kv.k), KLayout{}); + auto score = thread_qk.make_fragment_C(score_coordinates); + cute::clear(score); + auto q_source = thread_qk.partition_A(shared_q); + auto k_source = thread_qk.partition_B(shared_k_slot); + auto q_fragment = thread_qk.make_fragment_A(q_source); + auto k_fragment = thread_qk.make_fragment_B(k_source); + cute::warpgroup_fence_operand(score); + cute::warpgroup_arrive(); + cute::gemm(qk_mma, q_fragment, k_fragment, score); + cute::warpgroup_commit_batch(); + cute::warpgroup_wait<0>(); + cute::warpgroup_fence_operand(score); + + Array page_max; + page_max[0] = -CUDART_INF_F; + page_max[1] = -CUDART_INF_F; + CUTE_UNROLL + for (int i = 0; i < cute::size(score); ++i) { + const auto rk = score_coordinates(i); + const int row = cute::get<0>(rk); + const int key_in_page = cute::get<1>(rk); + const auto qh = decode_row(row); + const int query_position = query_base + cute::get<0>(qh); + const int head_in_group = cute::get<1>(qh); + const int absolute_key = key_begin + key_in_page; + const int last_valid = history_length + query_position; + const int first_valid = max(0, last_valid - arguments.window_size + 1); + const bool valid = query_position < query_length + && head_in_group < query_group_size + && absolute_key >= first_valid + && absolute_key <= last_valid; + score(i) = valid ? score(i) * arguments.qk_scale_log2 : + -CUDART_INF_F; + const int row_slot = row == row0 ? 0 : 1; + page_max[row_slot] = fmaxf(page_max[row_slot], score(i)); + } + CUTE_UNROLL + for (int offset = 1; offset <= 2; offset *= 2) { + page_max[0] = fmaxf( + page_max[0], __shfl_xor_sync(0xffffffffu, page_max[0], offset)); + page_max[1] = fmaxf( + page_max[1], __shfl_xor_sync(0xffffffffu, page_max[1], offset)); + } + CUTE_UNROLL + for (int row_slot = 0; row_slot < 2; ++row_slot) { + const float new_max = fmaxf(running_max[row_slot], page_max[row_slot]); + old_scale[row_slot] = running_max[row_slot] == -CUDART_INF_F ? + 0.f : exp2f(running_max[row_slot] - new_max); + running_max[row_slot] = new_max; + } + + Array page_sum; + page_sum[0] = 0.f; + page_sum[1] = 0.f; + CUTE_UNROLL + for (int i = 0; i < cute::size(score); ++i) { + const int row = cute::get<0>(score_coordinates(i)); + const int row_slot = row == row0 ? 0 : 1; + const float value = score(i); + const float probability_value = value == -CUDART_INF_F ? + 0.f : exp2f(value - running_max[row_slot]); + score(i) = probability_value; + page_sum[row_slot] += probability_value; + } + CUTE_UNROLL + for (int offset = 1; offset <= 2; offset *= 2) { + page_sum[0] += __shfl_xor_sync(0xffffffffu, page_sum[0], offset); + page_sum[1] += __shfl_xor_sync(0xffffffffu, page_sum[1], offset); + } + running_sum[0] = running_sum[0] * old_scale[0] + page_sum[0]; + running_sum[1] = running_sum[1] * old_scale[1] + page_sum[1]; + + CUTE_UNROLL + for (int i = 0; i < cute::size(output); ++i) { + const int row = cute::get<0>(pv_coordinates(i)); + output(i) *= row == row0 ? old_scale[0] : old_scale[1]; + } + CUTE_UNROLL + for (int i = 0; i < cute::size(score); ++i) { + probability(i) = static_cast(score(i)); + } + + if (slot == 0) { + cute::cp_async_wait<0>(); + __syncthreads(); + } + else if (has_next_pair) { + __syncthreads(); + if (next_page < last_page) { + auto next_k = cute::make_tensor( + cute::make_smem_ptr(storage.group[warp_group].body.kv.k), KLayout{}); + copy_paged_page( + cache, + next_page, + next_page * Policy::KeyTile, + key_length, + next_k, + local_tid); + cute::cp_async_fence(); + } + } + + auto shared_v_slot = cute::make_tensor( + cute::make_smem_ptr(storage.group[slot].body.kv.v), VLayout{}); + auto value_source = thread_pv.partition_B(shared_v_slot); + auto value_fragment = thread_pv.make_fragment_B(value_source); + cute::warpgroup_fence_operand(probability); + cute::warpgroup_fence_operand(output); + cute::warpgroup_arrive(); + cute::gemm(pv_mma, probability, value_fragment, output); + cute::warpgroup_commit_batch(); + cute::warpgroup_wait<0>(); + cute::warpgroup_fence_operand(probability); + cute::warpgroup_fence_operand(output); + } + } + + if (has_next_pair) { + __syncthreads(); + if (next_page < last_page) { + auto next_v = cute::make_tensor( + cute::make_smem_ptr(storage.group[warp_group].body.kv.v), VLayout{}); + auto next_v_copy = cute::composition( + next_v, + cute::Layout< + cute::Shape, + cute::Stride, cute::_1>>{}); + copy_paged_page( + cache, + next_page, + next_page * Policy::KeyTile, + key_length, + next_v_copy, + local_tid); + cute::cp_async_fence(); + } + } + } + } + } + else if (!inactive) { + PagedKv cache(arguments.block_ptrs, + arguments.block_ptr_offsets, + request, + kv_head, + arguments.kv_head_count, + arguments.block_len, + arguments.block_len_divmod, + arguments.cache_block_offset); + auto qk_mma = typename Mma::QK{}; + auto thread_qk = qk_mma.get_slice(local_tid); + auto qk_identity = cute::make_identity_tensor(cute::Shape{}); + auto score_coordinates = thread_qk.partition_C(qk_identity); + if (!QueryPartition || warp_group == 0) { + copy_paged_page( + cache, first_page, first_page * Policy::KeyTile, key_length, shared_k, local_tid); + cute::cp_async_fence(); + copy_paged_page( + cache, first_page, first_page * Policy::KeyTile, key_length, v_copy_view, local_tid); + cute::cp_async_fence(); + } + + for (int page = first_page; page < last_page; ++page) { + const int key_begin = page * Policy::KeyTile; + if constexpr (QueryPartition) { + if (warp_group == 0) { + cute::cp_async_wait<1>(); + } + __syncthreads(); + } + else { + cute::cp_async_wait<1>(); + sync_warp_group_barrier(warp_group); + } + + auto score = thread_qk.make_fragment_C(score_coordinates); + cute::clear(score); + auto q_source = thread_qk.partition_A(shared_q); + auto k_source = thread_qk.partition_B(shared_k); + auto q_fragment = thread_qk.make_fragment_A(q_source); + auto k_fragment = thread_qk.make_fragment_B(k_source); + cute::warpgroup_fence_operand(score); + cute::warpgroup_arrive(); + cute::gemm(qk_mma, q_fragment, k_fragment, score); + cute::warpgroup_commit_batch(); + cute::warpgroup_wait<0>(); + cute::warpgroup_fence_operand(score); + if constexpr (QueryPartition) { + __syncthreads(); + } + + const bool has_next_page = page + 1 < last_page; + if (has_next_page && (!QueryPartition || warp_group == 0)) { + const int next_page = page + 1; + copy_paged_page( + cache, + next_page, + next_page * Policy::KeyTile, + key_length, + shared_k, + local_tid); + cute::cp_async_fence(); + } + + Array page_max; + page_max[0] = -CUDART_INF_F; + page_max[1] = -CUDART_INF_F; + CUTE_UNROLL + for (int i = 0; i < cute::size(score); ++i) { + const auto rk = score_coordinates(i); + const int row = cute::get<0>(rk); + const int key_in_page = cute::get<1>(rk); + const auto qh = decode_row(row); + const int query_position = query_base + cute::get<0>(qh); + const int head_in_group = cute::get<1>(qh); + const int absolute_key = key_begin + key_in_page; + const int last_valid = history_length + query_position; + const int first_valid = max(0, last_valid - arguments.window_size + 1); + const bool valid = query_position < query_length + && head_in_group < query_group_size + && absolute_key >= first_valid + && absolute_key <= last_valid; + score(i) = valid ? score(i) * arguments.qk_scale_log2 : + -CUDART_INF_F; + const int row_slot = row == row0 ? 0 : 1; + page_max[row_slot] = fmaxf(page_max[row_slot], score(i)); + } + CUTE_UNROLL + for (int offset = 1; offset <= 2; offset *= 2) { + page_max[0] = fmaxf(page_max[0], + __shfl_xor_sync(0xffffffffu, page_max[0], offset)); + page_max[1] = fmaxf(page_max[1], + __shfl_xor_sync(0xffffffffu, page_max[1], offset)); + } + CUTE_UNROLL + for (int row_slot = 0; row_slot < 2; ++row_slot) { + const float new_max = fmaxf(running_max[row_slot], page_max[row_slot]); + old_scale[row_slot] = running_max[row_slot] == -CUDART_INF_F ? + 0.f : exp2f(running_max[row_slot] - new_max); + running_max[row_slot] = new_max; + } + + Array page_sum; + page_sum[0] = 0.f; + page_sum[1] = 0.f; + CUTE_UNROLL + for (int i = 0; i < cute::size(score); ++i) { + const int row = cute::get<0>(score_coordinates(i)); + const int row_slot = row == row0 ? 0 : 1; + const float value = score(i); + const float probability = value == -CUDART_INF_F ? + 0.f : exp2f(value - running_max[row_slot]); + score(i) = probability; + page_sum[row_slot] += probability; + } + CUTE_UNROLL + for (int offset = 1; offset <= 2; offset *= 2) { + page_sum[0] += __shfl_xor_sync(0xffffffffu, page_sum[0], offset); + page_sum[1] += __shfl_xor_sync(0xffffffffu, page_sum[1], offset); + } + running_sum[0] = running_sum[0] * old_scale[0] + page_sum[0]; + running_sum[1] = running_sum[1] * old_scale[1] + page_sum[1]; + + CUTE_UNROLL + for (int i = 0; i < cute::size(output); ++i) { + const int row = cute::get<0>(pv_coordinates(i)); + output(i) *= row == row0 ? old_scale[0] : old_scale[1]; + } + CUTE_UNROLL + for (int i = 0; i < cute::size(score); ++i) { + probability(i) = static_cast(score(i)); + } + if constexpr (QueryPartition) { + if (warp_group == 0) { + if (has_next_page) { + cute::cp_async_wait<1>(); + } + else { + cute::cp_async_wait<0>(); + } + } + __syncthreads(); + } + else { + if (has_next_page) { + cute::cp_async_wait<1>(); + } + else { + cute::cp_async_wait<0>(); + } + sync_warp_group_barrier(warp_group); + } + + auto value_source = thread_pv.partition_B(shared_v); + auto value_fragment = thread_pv.make_fragment_B(value_source); + cute::warpgroup_fence_operand(probability); + cute::warpgroup_fence_operand(output); + cute::warpgroup_arrive(); + cute::gemm(pv_mma, probability, value_fragment, output); + cute::warpgroup_commit_batch(); + cute::warpgroup_wait<0>(); + cute::warpgroup_fence_operand(probability); + cute::warpgroup_fence_operand(output); + if constexpr (QueryPartition) { + __syncthreads(); + } + + if (has_next_page && (!QueryPartition || warp_group == 0)) { + const int next_page = page + 1; + copy_paged_page( + cache, + next_page, + next_page * Policy::KeyTile, + key_length, + v_copy_view, + local_tid); + cute::cp_async_fence(); + } + } + } + + CUTE_UNROLL + for (int i = 0; i < cute::size(output); ++i) { + const auto rd = pv_coordinates(i); + const int row = cute::get<0>(rd); + const auto qh = decode_row(row); + const int query_position = query_base + cute::get<0>(qh); + const int head_in_group = cute::get<1>(qh); + if (query_position < query_length && head_in_group < query_group_size) { + const int absolute_query = query_begin + query_position; + const int query_head = kv_head * query_group_size + head_in_group; + const int d = cute::get<1>(rd); + if constexpr (StorePartial) { + const int local_query = absolute_query - arguments.query_offset; + partial_o(local_query, split, query_head, d) = output(i); + } + else if constexpr (QueryPartition) { + const int row_slot = row == row0 ? 0 : 1; + const float sum = running_sum[row_slot]; + global_out(absolute_query, query_head, d) = + sum == 0.f ? T(0) : static_cast(output(i) / sum); + } + else { + auto merge_output = cute::make_tensor( + cute::make_smem_ptr(group_storage.body.output), + cute::make_layout(cute::Shape{}, + cute::Stride{})); + merge_output(row, d) = output(i); + } + } + } + + const int lane = local_tid % 32; + if ((!QueryPartition || StorePartial) && lane % 4 == 0) { + group_storage.running_max[row0] = running_max[0]; + group_storage.running_max[row1] = running_max[1]; + group_storage.running_sum[row0] = running_sum[0]; + group_storage.running_sum[row1] = running_sum[1]; + } + + if constexpr (StorePartial) { + if (lane % 4 == 0) { + CUTE_UNROLL + for (int row_slot = 0; row_slot < 2; ++row_slot) { + const int row = row_slot == 0 ? row0 : row1; + const auto qh = decode_row(row); + const int query_position = query_base + cute::get<0>(qh); + const int head_in_group = cute::get<1>(qh); + if (query_position < query_length && head_in_group < query_group_size) { + const int absolute_query = query_begin + query_position; + const int local_query = absolute_query - arguments.query_offset; + const int query_head = kv_head * query_group_size + head_in_group; + partial_ml(local_query, split, query_head, cute::_0{}) = running_max[row_slot]; + partial_ml(local_query, split, query_head, cute::_1{}) = running_sum[row_slot]; + } + } + } + } + else if constexpr (!QueryPartition) { + __syncthreads(); + if (warp_group != 0) { + return; + } + auto output0 = cute::make_tensor( + cute::make_smem_ptr(storage.group[0].body.output), + cute::make_layout(cute::Shape{}, + cute::Stride{})); + auto output1 = cute::make_tensor( + cute::make_smem_ptr(storage.group[1].body.output), + cute::make_layout(cute::Shape{}, + cute::Stride{})); + CUTE_UNROLL + for (int access = 0; access < Policy::MTile * Policy::HeadDim / Policy::Threads; ++access) { + const int linear = local_tid + access * Policy::Threads; + const int row = linear / Policy::HeadDim; + const int d = linear % Policy::HeadDim; + const auto qh = decode_row(row); + const int query_position = cute::get<0>(qh); + const int head_in_group = cute::get<1>(qh); + if (query_position < query_length && head_in_group < query_group_size) { + const float max0 = storage.group[0].running_max[row]; + const float max1 = storage.group[1].running_max[row]; + const float merged_max = fmaxf(max0, max1); + const float scale0 = max0 == -CUDART_INF_F ? 0.f : exp2f(max0 - merged_max); + const float scale1 = max1 == -CUDART_INF_F ? 0.f : exp2f(max1 - merged_max); + const float sum0 = storage.group[0].running_sum[row]; + const float sum1 = storage.group[1].running_sum[row]; + const float merged_sum = sum0 * scale0 + sum1 * scale1; + const float inverse_sum = merged_sum == 0.f ? 0.f : 1.f / merged_sum; + const float value = output0(row, d) * scale0 + output1(row, d) * scale1; + const int absolute_query = query_begin + query_position; + const int query_head = kv_head * query_group_size + head_in_group; + global_out(absolute_query, query_head, d) = static_cast(value * inverse_sum); + } + } + } + } +}; + +template +__global__ __launch_bounds__(256, 1) +void VerificationAttentionWgmmaRsKernel(Arguments arguments) +{ + extern __shared__ char dynamic_shared[]; + auto& storage = *reinterpret_cast*>(dynamic_shared); + Sm90WgmmaRsMainloop{arguments, storage}.run(); +} + +template +void LaunchWgmmaRs(const Arguments& arguments) +{ + using Policy = Sm90WgmmaRsPolicy256; + const dim3 grid(arguments.request_count, + arguments.kv_head_count, + (arguments.split_count + 1) / 2); + constexpr int smem_bytes = sizeof(Sm90WgmmaRsSharedStorage); + if (arguments.split_count == 1) { + auto kernel = VerificationAttentionWgmmaRsKernel; + cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_bytes); + cudaFuncSetAttribute(kernel, cudaFuncAttributePreferredSharedMemoryCarveout, 100); + kernel<<>>(arguments); + } + else { + auto kernel = VerificationAttentionWgmmaRsKernel; + cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_bytes); + cudaFuncSetAttribute(kernel, cudaFuncAttributePreferredSharedMemoryCarveout, 100); + kernel<<>>(arguments); + } +} + +} // namespace turbomind::verification_attention diff --git a/src/turbomind/kernels/attention/verification/kernel_sm90_wgmma_rs_ws.cuh b/src/turbomind/kernels/attention/verification/kernel_sm90_wgmma_rs_ws.cuh new file mode 100644 index 0000000000..2d99f2f694 --- /dev/null +++ b/src/turbomind/kernels/attention/verification/kernel_sm90_wgmma_rs_ws.cuh @@ -0,0 +1,790 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#pragma once + +#include + +#include +#include + +#include "src/turbomind/kernels/attention/verification/kernel_sm90_wgmma_rs.cuh" + +namespace turbomind::verification_attention { + +struct Sm90WgmmaRsWsComputePolicy { + static constexpr int HeadDim = 256; + static constexpr int Threads = 128; + static constexpr int MTile = 64; + static constexpr int KeyTile = 64; + + using QShape = cute::Shape; +}; + +struct Sm90WgmmaRsWsCopyPolicy { + static constexpr int HeadDim = 256; + static constexpr int Threads = 128; + static constexpr int KeyTile = 64; + + using KvShape = cute::Shape; +}; + +template +struct Sm90WgmmaRsWsMma; + +template<> +struct Sm90WgmmaRsWsMma { + using QK = decltype(cute::make_tiled_mma( + cute::SM90_64x64x16_F32BF16BF16_SS< + cute::GMMA::Major::K, cute::GMMA::Major::K>{})); + using PV = decltype(cute::make_tiled_mma( + cute::SM90_64x256x16_F32BF16BF16_RS< + cute::GMMA::Major::K, cute::GMMA::Major::MN>{})); +}; + +using Sm90WgmmaRsWsKPipeline = cutlass::PipelineAsync<2>; +using Sm90WgmmaRsWsVPipeline = cutlass::PipelineAsync<2>; + +template +struct alignas(128) Sm90WgmmaRsWsSharedStorage { + using ComputePolicy = Sm90WgmmaRsWsComputePolicy; + using CopyPolicy = Sm90WgmmaRsWsCopyPolicy; + + alignas(128) T q[2][ComputePolicy::MTile * ComputePolicy::HeadDim]; + alignas(128) T k[2][CopyPolicy::KeyTile * CopyPolicy::HeadDim]; + alignas(128) T v[2][CopyPolicy::KeyTile * CopyPolicy::HeadDim]; + alignas(16) typename Sm90WgmmaRsWsKPipeline::SharedStorage k_pipeline; + alignas(16) typename Sm90WgmmaRsWsVPipeline::SharedStorage v_pipeline; +}; + +static_assert(sizeof(Sm90WgmmaRsWsSharedStorage) <= 232448); + +template +struct Sm90WgmmaRsWsMainloop { + using ComputePolicy = Sm90WgmmaRsWsComputePolicy; + using CopyPolicy = Sm90WgmmaRsWsCopyPolicy; + using Mma = Sm90WgmmaRsWsMma; + using KPipeline = Sm90WgmmaRsWsKPipeline; + using VPipeline = Sm90WgmmaRsWsVPipeline; + using Storage = Sm90WgmmaRsWsSharedStorage; + + Arguments arguments; + Storage& storage; + + CUTE_DEVICE auto decode_row(int compute_group, int row) const + { + const int flat_row = compute_group * ComputePolicy::MTile + row; + int query_position; + int head_in_group; + arguments.query_group_size_divmod(query_position, head_in_group, flat_row); + return cute::make_coord(query_position, head_in_group); + } + + CUTE_DEVICE auto make_k_raw() + { + using Layout = decltype(cute::tile_to_shape( + cute::GMMA::Layout_K_SW128_Atom{}, + cute::Shape{})); + return cute::make_tensor(cute::make_smem_ptr(&storage.k[0][0]), Layout{}); + } + + CUTE_DEVICE auto make_v_raw() + { + using Layout = decltype(cute::tile_to_shape( + cute::GMMA::Layout_MN_SW128_Atom{}, + cute::Shape{})); + return cute::make_tensor(cute::make_smem_ptr(&storage.v[0][0]), Layout{}); + } + + CUTE_DEVICE void producer(KPipeline& k_pipeline, + VPipeline& v_pipeline, + int request, + int kv_head, + int key_length, + int first_page, + int last_page, + bool inactive) + { + cutlass::arch::warpgroup_reg_dealloc<40>(); + if (inactive) { + return; + } + + PagedKv cache(arguments.block_ptrs, + arguments.block_ptr_offsets, + request, + kv_head, + arguments.kv_head_count, + arguments.block_len, + arguments.block_len_divmod, + arguments.cache_block_offset); + auto shared_k_raw = make_k_raw(); + auto shared_v_raw = make_v_raw(); + const int local_tid = threadIdx.x; + auto k_write = cutlass::make_producer_start_state(); + auto v_write = cutlass::make_producer_start_state(); + const int full_page_count = key_length / CopyPolicy::KeyTile; + auto load_k = [&](int page) { + k_pipeline.producer_acquire(k_write); + auto destination = shared_k_raw(cute::_, cute::_, k_write.index()); + if (page < full_page_count) { + copy_paged_page_full( + cache, page, destination, local_tid); + } + else { + copy_paged_page( + cache, + page, + page * CopyPolicy::KeyTile, + key_length, + destination, + local_tid); + } + k_pipeline.producer_commit(k_write, cutlass::arch::cpasync_barrier_arrive); + ++k_write; + }; + auto load_v = [&](int page) { + v_pipeline.producer_acquire(v_write); + auto destination = cute::composition( + shared_v_raw(cute::_, cute::_, v_write.index()), + cute::Layout< + cute::Shape, + cute::Stride>{}); + if (page < full_page_count) { + copy_paged_page_full( + cache, page, destination, local_tid); + } + else { + copy_paged_page( + cache, + page, + page * CopyPolicy::KeyTile, + key_length, + destination, + local_tid); + } + v_pipeline.producer_commit(v_write, cutlass::arch::cpasync_barrier_arrive); + ++v_write; + }; + for (int page = first_page; page < last_page; ++page) { + load_k(page); + load_v(page); + } + if (local_tid == 0) { + k_pipeline.producer_tail(k_write); + v_pipeline.producer_tail(v_write); + } + } + + CUTE_DEVICE void load_query(int compute_group, + int local_tid, + int request, + int kv_head, + int query_begin, + int query_length, + int history_length) + { + using QueryCopy = QTileCopy; + using QLayout = decltype(cute::tile_to_shape( + cute::GMMA::Layout_K_SW128_Atom{}, + cute::Shape{})); + auto shared_q = cute::make_tensor( + cute::make_smem_ptr(storage.q[compute_group]), QLayout{}); + auto global_q = cute::make_tensor( + cute::make_gmem_ptr(static_cast(arguments.q)), + cute::make_layout( + cute::make_shape(arguments.query_count, + arguments.query_head_count, + cute::_256{}), + cute::make_stride(arguments.q_stride, + cute::Int{}, + cute::_1{}))); + auto global_q_bias = cute::make_tensor( + cute::make_gmem_ptr(static_cast(arguments.q_bias)), + cute::make_layout( + cute::make_shape(arguments.query_head_count, cute::_256{}), + cute::make_stride(cute::Int{}, cute::_1{}))); + using RegisterCopy = cute::Copy_Atom, T>; + auto q_identity = cute::make_identity_tensor(typename ComputePolicy::QShape{}); + auto tiled_q_copy = typename QueryCopy::TiledCopy{}; + auto thread_q_copy = tiled_q_copy.get_thread_slice(local_tid); + auto q_coordinates = cute::group_modes< + 1, cute::rank_v>( + thread_q_copy.partition_S(q_identity)); + auto q_destinations = cute::group_modes< + 1, cute::rank_v>( + thread_q_copy.partition_D(shared_q)); + Array fragment; + auto fragment_tensor = cute::make_tensor( + cute::make_rmem_ptr(fragment.data()), + cute::make_layout(cute::Int{})); + Array bias; + auto bias_fragment = cute::make_tensor( + cute::make_rmem_ptr(bias.data()), + cute::make_layout(cute::Int{})); + const int query_group_size = arguments.query_group_size; + const int d_begin = cute::get<1>(q_coordinates(cute::_0{}, cute::_0{})); + FastRoPE rope( + arguments.rope, + request, + std::integral_constant{}); + const bool rotary = d_begin < arguments.rope.dim; + if (rotary) { + rope.init(d_begin); + } + const auto first_qh = decode_row( + compute_group, cute::get<0>(q_coordinates(cute::_0{}, cute::_0{}))); + int query_position = cute::get<0>(first_qh); + int head_in_group = cute::get<1>(first_qh); + + CUTE_UNROLL + for (int access = 0; access < QueryCopy::AccessCount; ++access) { + const bool valid = query_position < query_length; + CUTE_UNROLL + for (int value = 0; value < QueryCopy::ValuesPerAccess; ++value) { + fragment[value] = T(0); + } + if (valid) { + const int query_head = kv_head * query_group_size + head_in_group; + auto source = cute::make_tensor( + cute::make_gmem_ptr(&global_q(query_begin + query_position, + query_head, + d_begin)), + cute::make_layout(cute::Int{})); + cute::copy(RegisterCopy{}, source, fragment_tensor); + if (arguments.q_bias) { + auto bias_source = cute::make_tensor( + cute::make_gmem_ptr(&global_q_bias(query_head, d_begin)), + cute::make_layout(cute::Int{})); + cute::copy(RegisterCopy{}, bias_source, bias_fragment); + CUTE_UNROLL + for (int value = 0; value < QueryCopy::ValuesPerAccess; ++value) { + fragment[value] = fragment[value] + bias[value]; + } + } + if (rotary) { + rope.apply(fragment, history_length + query_position); + } + } + cute::copy(RegisterCopy{}, fragment_tensor, q_destinations(cute::_, access)); + head_in_group += QueryCopy::RowsPerThreadTile; + if (head_in_group >= query_group_size) { + head_in_group -= query_group_size; + ++query_position; + } + } + cutlass::arch::fence_view_async_shared(); + sync_warp_group_barrier(compute_group); + } + + CUTE_DEVICE void consumer(KPipeline& k_pipeline, + VPipeline& v_pipeline, + int request, + int kv_head, + int query_begin, + int query_length, + int key_length, + int history_length, + int first_page, + int last_page, + int split, + bool inactive) + { + cutlass::arch::warpgroup_reg_alloc<232>(); + const int compute_group = threadIdx.x / 128 - 1; + const int local_tid = threadIdx.x % 128; + const int query_group_size = arguments.query_group_size; + if (!inactive) { + load_query(compute_group, + local_tid, + request, + kv_head, + query_begin, + query_length, + history_length); + } + + using QLayout = decltype(cute::tile_to_shape( + cute::GMMA::Layout_K_SW128_Atom{}, + cute::Shape{})); + auto shared_q = cute::make_tensor( + cute::make_smem_ptr(storage.q[compute_group]), QLayout{}); + auto shared_k = make_k_raw(); + auto shared_v = make_v_raw(); + auto qk_mma = typename Mma::QK{}; + auto pv_mma = typename Mma::PV{}; + auto thread_qk = qk_mma.get_slice(local_tid); + auto thread_pv = pv_mma.get_slice(local_tid); + auto qk_identity = cute::make_identity_tensor(cute::Shape{}); + auto pv_identity = cute::make_identity_tensor(cute::Shape{}); + auto score_coordinates = thread_qk.partition_C(qk_identity); + auto pv_coordinates = thread_pv.partition_C(pv_identity); + const int row0 = cute::get<0>(pv_coordinates(0)); + const int row1 = cute::get<0>(pv_coordinates(2)); + const auto qh0 = decode_row(compute_group, row0); + const auto qh1 = decode_row(compute_group, row1); + Array query_position; + Array head_in_group; + query_position[0] = cute::get<0>(qh0); + query_position[1] = cute::get<0>(qh1); + head_in_group[0] = cute::get<1>(qh0); + head_in_group[1] = cute::get<1>(qh1); + auto output = thread_pv.make_fragment_C(pv_coordinates); + cute::clear(output); + + using ScoreTensor = decltype(thread_qk.make_fragment_C(score_coordinates)); + using ScoreLayout = typename ScoreTensor::layout_type; + static_assert(cute::cosize_v == 32); + using ProbabilityLayout = cute::Layout< + cute::Shape, cute::_1, cute::_4>, + cute::Stride, cute::_0, cute::_8>>; + cutlass::Array score_storage; + auto score = cute::make_tensor( + cute::make_rmem_ptr(score_storage.data()), ScoreLayout{}); + cutlass::Array probability_storage; + auto probability = cute::make_tensor( + cute::make_rmem_ptr(probability_storage.data()), ProbabilityLayout{}); + auto output_coordinates = pv_coordinates; + + Array running_max; + Array running_sum; + Array old_scale; + CUTE_UNROLL + for (int row_slot = 0; row_slot < 2; ++row_slot) { + running_max[row_slot] = -CUDART_INF_F; + running_sum[row_slot] = 0.f; + old_scale[row_slot] = 0.f; + } + + typename KPipeline::PipelineState k_read; + typename VPipeline::PipelineState v_read; + if (!inactive) { + const int first_flat_row = compute_group * ComputePolicy::MTile; + const int valid_flat_row_end = min(first_flat_row + ComputePolicy::MTile, + query_length * query_group_size); + const int first_query_position = + arguments.query_group_size_divmod.div(first_flat_row); + const int last_query_position = + arguments.query_group_size_divmod.div(valid_flat_row_end - 1); + const bool full_query_tile = first_flat_row < valid_flat_row_end; + const int full_page_first_key = + max(0, history_length + last_query_position - arguments.window_size + 1); + const int full_page_last_key = history_length + first_query_position; + + auto q_source = thread_qk.partition_A(shared_q); + auto q_fragment = thread_qk.make_fragment_A(q_source); + + auto issue_qk = [&] { + cute::clear(score); + auto k_source = thread_qk.partition_B( + shared_k(cute::_, cute::_, k_read.index())); + auto k_fragment = thread_qk.make_fragment_B(k_source); + cute::warpgroup_fence_operand(score); + cute::warpgroup_arrive(); + cute::gemm(qk_mma, q_fragment, k_fragment, score); + cute::warpgroup_commit_batch(); + }; + + auto update_softmax = [&](int page, auto full_page_tag) { + constexpr bool FullPage = decltype(full_page_tag)::value; + const int key_begin = page * CopyPolicy::KeyTile; + Array page_max; + page_max[0] = -CUDART_INF_F; + page_max[1] = -CUDART_INF_F; + if constexpr (FullPage) { + CUTE_UNROLL + for (int i = 0; i < cute::size(score); ++i) { + const int row = cute::get<0>(score_coordinates(i)); + const int row_slot = row == row0 ? 0 : 1; + page_max[row_slot] = fmaxf(page_max[row_slot], score(i)); + } + } + else { + CUTE_UNROLL + for (int i = 0; i < cute::size(score); ++i) { + const auto rk = score_coordinates(i); + const int row = cute::get<0>(rk); + const int key_in_page = cute::get<1>(rk); + const int row_slot = row == row0 ? 0 : 1; + const int absolute_key = key_begin + key_in_page; + const int last_valid = history_length + query_position[row_slot]; + const int first_valid = max(0, last_valid - arguments.window_size + 1); + const bool valid = query_position[row_slot] < query_length + && absolute_key >= first_valid + && absolute_key <= last_valid; + score(i) = valid ? score(i) : -CUDART_INF_F; + page_max[row_slot] = fmaxf(page_max[row_slot], score(i)); + } + } + CUTE_UNROLL + for (int offset = 1; offset <= 2; offset *= 2) { + page_max[0] = fmaxf( + page_max[0], __shfl_xor_sync(0xffffffffu, page_max[0], offset)); + page_max[1] = fmaxf( + page_max[1], __shfl_xor_sync(0xffffffffu, page_max[1], offset)); + } + CUTE_UNROLL + for (int row_slot = 0; row_slot < 2; ++row_slot) { + page_max[row_slot] *= arguments.qk_scale_log2; + const float new_max = fmaxf(running_max[row_slot], page_max[row_slot]); + old_scale[row_slot] = running_max[row_slot] == -CUDART_INF_F ? + 0.f : exp2f(running_max[row_slot] - new_max); + running_max[row_slot] = new_max; + } + + Array page_sum; + page_sum[0] = 0.f; + page_sum[1] = 0.f; + if constexpr (FullPage) { + CUTE_UNROLL + for (int i = 0; i < cute::size(score); ++i) { + const int row = cute::get<0>(score_coordinates(i)); + const int row_slot = row == row0 ? 0 : 1; + const float p = exp2f( + fmaf(score(i), arguments.qk_scale_log2, -running_max[row_slot])); + score(i) = p; + page_sum[row_slot] += p; + } + } + else { + CUTE_UNROLL + for (int i = 0; i < cute::size(score); ++i) { + const int row = cute::get<0>(score_coordinates(i)); + const int row_slot = row == row0 ? 0 : 1; + const float value = score(i); + const float p = value == -CUDART_INF_F ? + 0.f : exp2f(fmaf(value, + arguments.qk_scale_log2, + -running_max[row_slot])); + score(i) = p; + page_sum[row_slot] += p; + } + } + CUTE_UNROLL + for (int offset = 1; offset <= 2; offset *= 2) { + page_sum[0] += __shfl_xor_sync(0xffffffffu, page_sum[0], offset); + page_sum[1] += __shfl_xor_sync(0xffffffffu, page_sum[1], offset); + } + running_sum[0] = fmaf(running_sum[0], old_scale[0], page_sum[0]); + running_sum[1] = fmaf(running_sum[1], old_scale[1], page_sum[1]); + CUTE_UNROLL + for (int i = 0; i < cute::size(probability); ++i) { + probability(i) = static_cast(score(i)); + } + }; + + auto scale_output = [&] { + const bool warp_needs_scale = __any_sync( + 0xffffffffu, old_scale[0] != 1.f || old_scale[1] != 1.f); + if (!warp_needs_scale) { + return; + } + CUTE_UNROLL + for (int i = 0; i < cute::size(output); ++i) { + const int row = cute::get<0>(output_coordinates(i)); + output(i) *= row == row0 ? old_scale[0] : old_scale[1]; + } + }; + + auto issue_pv = [&] { + auto value_source = thread_pv.partition_B( + shared_v(cute::_, cute::_, v_read.index())); + auto value_fragment = thread_pv.make_fragment_B(value_source); + cute::warpgroup_fence_operand(probability); + cute::warpgroup_fence_operand(output); + cute::warpgroup_arrive(); + cute::gemm(pv_mma, probability, value_fragment, output); + cute::warpgroup_commit_batch(); + }; + + auto process_page = [&](int page, auto full_page_tag) { + k_pipeline.consumer_wait(k_read); + issue_qk(); + auto v_ready = v_pipeline.consumer_try_wait(v_read); + cute::warpgroup_wait<0>(); + cute::warpgroup_fence_operand(score); + k_pipeline.consumer_release(k_read); + ++k_read; + update_softmax(page, full_page_tag); + scale_output(); + v_pipeline.consumer_wait(v_read, v_ready); + issue_pv(); + cute::warpgroup_wait<0>(); + cute::warpgroup_fence_operand(probability); + cute::warpgroup_fence_operand(output); + v_pipeline.consumer_release(v_read); + ++v_read; + }; + + const int full_page_begin = full_query_tile ? + min(last_page, + max(first_page, + (full_page_first_key + CopyPolicy::KeyTile - 1) / + CopyPolicy::KeyTile)) : last_page; + const int full_page_end = full_query_tile ? + max(full_page_begin, + min(last_page, + (full_page_last_key + 1) / CopyPolicy::KeyTile)) : last_page; + for (int page = first_page; page < full_page_begin; ++page) { + process_page(page, std::false_type{}); + } + for (int page = full_page_begin; page < full_page_end; ++page) { + process_page(page, std::true_type{}); + } + for (int page = full_page_end; page < last_page; ++page) { + process_page(page, std::false_type{}); + } + } + + auto global_out = cute::make_tensor( + cute::make_gmem_ptr(static_cast(arguments.out)), + cute::make_layout( + cute::make_shape(arguments.query_count, + arguments.query_head_count, + cute::_256{}), + cute::make_stride(arguments.query_head_count * ComputePolicy::HeadDim, + cute::Int{}, + cute::_1{}))); + auto partial_o = cute::make_tensor( + cute::make_gmem_ptr(arguments.partial_o), + cute::make_layout( + cute::make_shape(arguments.query_count, + arguments.split_count, + arguments.query_head_count, + cute::_256{}), + cute::make_stride(arguments.split_count * arguments.query_head_count * ComputePolicy::HeadDim, + arguments.query_head_count * ComputePolicy::HeadDim, + cute::Int{}, + cute::_1{}))); + auto partial_ml = cute::make_tensor( + cute::make_gmem_ptr(arguments.partial_ml), + cute::make_layout( + cute::make_shape(arguments.query_count, + arguments.split_count, + arguments.query_head_count, + cute::_2{}), + cute::make_stride(arguments.split_count * arguments.query_head_count * 2, + arguments.query_head_count * 2, + cute::_2{}, + cute::_1{}))); + if constexpr (StorePartial) { + CUTE_UNROLL + for (int i = 0; i < cute::size(output); ++i) { + const auto rd = output_coordinates(i); + const int row = cute::get<0>(rd); + const auto qh = decode_row(compute_group, row); + const int position = cute::get<0>(qh); + const int head = cute::get<1>(qh); + if (position < query_length) { + const int absolute_query = query_begin + position; + const int local_query = absolute_query - arguments.query_offset; + const int query_head = kv_head * query_group_size + head; + const int d = cute::get<1>(rd); + partial_o(local_query, split, query_head, d) = output(i); + } + } + } + else { + using OutputLayout = decltype(cute::tile_to_shape( + cute::GMMA::Layout_K_SW128_Atom{}, + cute::Shape{})); + using OutputStore = decltype(cute::make_tiled_copy_C( + cute::Copy_Atom{}, + typename Mma::PV{})); + static constexpr int OutputElementsPerStore = 16 / sizeof(T); + static constexpr int OutputThreadsPerRow = 64 / OutputElementsPerStore; + using OutputGlobalCopy = decltype(cute::make_tiled_copy( + cute::Copy_Atom, T>{}, + cute::Layout< + cute::Shape, + cute::Int>, + cute::Stride, cute::_1>>{}, + cute::Layout>>{})); + + cutlass::Array inverse_sum; + CUTE_UNROLL + for (int row_slot = 0; row_slot < 2; ++row_slot) { + const float sum = running_sum[row_slot]; + inverse_sum[row_slot] = sum == 0.f ? 0.f : 1.f / sum; + } + + if (query_length * query_group_size < 2 * ComputePolicy::MTile) { + CUTE_UNROLL + for (int i = 0; i < cute::size(output); ++i) { + const auto rd = output_coordinates(i); + const int row = cute::get<0>(rd); + const auto qh = decode_row(compute_group, row); + const int position = cute::get<0>(qh); + if (position < query_length) { + const int row_slot = row == row0 ? 0 : 1; + const int head = cute::get<1>(qh); + const int absolute_query = query_begin + position; + const int query_head = kv_head * query_group_size + head; + const int d = cute::get<1>(rd); + global_out(absolute_query, query_head, d) = static_cast( + output(i) * inverse_sum[row_slot]); + } + } + return; + } + + auto normalized = cute::make_tensor_like(output); + CUTE_UNROLL + for (int i = 0; i < cute::size(output); ++i) { + const int row = cute::get<0>(output_coordinates(i)); + const int row_slot = row == row0 ? 0 : 1; + normalized(i) = static_cast(output(i) * inverse_sum[row_slot]); + } + + // Both compute warp groups consumed the common K ring. Reuse two + // dead K stages as disjoint 64x256 epilogue tiles. + sync_compute_groups(); + auto shared_output = cute::make_tensor( + cute::make_smem_ptr(storage.k[compute_group]), OutputLayout{}); + auto output_store = OutputStore{}; + auto thread_output_store = output_store.get_thread_slice(local_tid); + cute::copy(output_store, + thread_output_store.retile_S(normalized), + thread_output_store.partition_D(shared_output)); + cutlass::arch::fence_view_async_shared(); + sync_warp_group_barrier(compute_group); + + auto output_identity = cute::make_identity_tensor( + cute::Shape{}); + auto output_copy = OutputGlobalCopy{}; + auto thread_output_copy = output_copy.get_thread_slice(local_tid); + auto shared_fragments = thread_output_copy.partition_S(shared_output); + auto output_coordinates_copy = thread_output_copy.partition_D(output_identity); + CUTE_UNROLL + for (int m = 0; m < cute::size<1>(shared_fragments); ++m) { + const int row = cute::get<0>(output_coordinates_copy(cute::_0{}, m, cute::_0{})); + const auto qh = decode_row(compute_group, row); + const int position = cute::get<0>(qh); + const int head = cute::get<1>(qh); + if (position < query_length) { + const int absolute_query = query_begin + position; + const int query_head = kv_head * query_group_size + head; + auto output_row = global_out(absolute_query, query_head, cute::_); + auto output_vectors = cute::tiled_divide( + output_row, + cute::Shape>{}); + CUTE_UNROLL + for (int k = 0; k < cute::size<2>(shared_fragments); ++k) { + const int d = cute::get<1>( + output_coordinates_copy(cute::_0{}, cute::_0{}, k)); + cute::copy(output_copy, + shared_fragments(cute::_, m, k), + output_vectors(cute::_, d / OutputElementsPerStore)); + } + } + } + } + + if constexpr (StorePartial) { + const int lane = local_tid % 32; + if (lane % 4 == 0) { + CUTE_UNROLL + for (int row_slot = 0; row_slot < 2; ++row_slot) { + if (query_position[row_slot] < query_length) { + const int absolute_query = query_begin + query_position[row_slot]; + const int local_query = absolute_query - arguments.query_offset; + const int query_head = + kv_head * query_group_size + head_in_group[row_slot]; + partial_ml(local_query, split, query_head, cute::_0{}) = running_max[row_slot]; + partial_ml(local_query, split, query_head, cute::_1{}) = running_sum[row_slot]; + } + } + } + } + } + + CUTE_DEVICE void run() + { + const int warp_group = threadIdx.x / 128; + typename KPipeline::Params k_params; + k_params.role = warp_group == 0 ? KPipeline::ThreadCategory::Producer : + KPipeline::ThreadCategory::Consumer; + k_params.producer_arv_count = 128; + k_params.consumer_arv_count = 256; + k_params.initializing_warp = 0; + typename VPipeline::Params v_params; + v_params.role = warp_group == 0 ? VPipeline::ThreadCategory::Producer : + VPipeline::ThreadCategory::Consumer; + v_params.producer_arv_count = 128; + v_params.consumer_arv_count = 256; + v_params.initializing_warp = 0; + KPipeline k_pipeline(storage.k_pipeline, k_params); + VPipeline v_pipeline(storage.v_pipeline, v_params); + __syncthreads(); + + const int request = blockIdx.x; + const int kv_head = blockIdx.y; + const int query_begin = arguments.q_offsets[request]; + const int query_end = arguments.q_offsets[request + 1]; + const int query_length = query_end - query_begin; + const int key_length = arguments.k_offsets[request + 1] - arguments.k_offsets[request]; + const int history_length = key_length - query_length; + const int page_count = (key_length + CopyPolicy::KeyTile - 1) / CopyPolicy::KeyTile; + const int split = blockIdx.z; + const int first_page = static_cast(page_count) * split / arguments.split_count; + const int last_page = static_cast(page_count) * (split + 1) / arguments.split_count; + const bool request_finished = arguments.finished && arguments.finished[request]; + const bool inactive = request_finished || first_page >= last_page; + + if (warp_group == 0) { + producer(k_pipeline, + v_pipeline, + request, + kv_head, + key_length, + first_page, + last_page, + inactive); + } + else { + consumer(k_pipeline, + v_pipeline, + request, + kv_head, + query_begin, + query_length, + key_length, + history_length, + first_page, + last_page, + split, + inactive); + } + } +}; + +template +__global__ __launch_bounds__(384, 1) +void VerificationAttentionWgmmaRsWsKernel(Arguments arguments) +{ + extern __shared__ char dynamic_shared[]; + auto& storage = *reinterpret_cast*>(dynamic_shared); + Sm90WgmmaRsWsMainloop{arguments, storage}.run(); +} + +template +void LaunchWgmmaRsWs(const Arguments& arguments) +{ + const dim3 grid(arguments.request_count, + arguments.kv_head_count, + arguments.split_count); + constexpr int smem_bytes = sizeof(Sm90WgmmaRsWsSharedStorage); + if (arguments.split_count == 1) { + auto kernel = VerificationAttentionWgmmaRsWsKernel; + cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_bytes); + cudaFuncSetAttribute(kernel, cudaFuncAttributePreferredSharedMemoryCarveout, 100); + kernel<<>>(arguments); + } + else { + auto kernel = VerificationAttentionWgmmaRsWsKernel; + cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_bytes); + cudaFuncSetAttribute(kernel, cudaFuncAttributePreferredSharedMemoryCarveout, 100); + kernel<<>>(arguments); + } +} + +} // namespace turbomind::verification_attention diff --git a/src/turbomind/kernels/attention/verification/paged_kv.cuh b/src/turbomind/kernels/attention/verification/paged_kv.cuh new file mode 100644 index 0000000000..dc3b552a12 --- /dev/null +++ b/src/turbomind/kernels/attention/verification/paged_kv.cuh @@ -0,0 +1,161 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#pragma once + +#include + +#include "src/turbomind/kernels/attention/block.h" +#include "src/turbomind/kernels/attention/verification/policy_sm90.cuh" + +namespace turbomind::verification_attention { + +template +class PagedKv { +public: + using Config = block::Config; + using Layout = block::Layout; + + CUTE_DEVICE PagedKv(char* const* block_ptrs, + const int* block_ptr_offsets, + int request, + int kv_head, + int kv_head_count, + int block_len, + cutlass::FastDivmod block_len_divmod, + int cache_block_offset): + pages_{block_ptrs + block_ptr_offsets[request]}, + layout_{Config{kv_head_count, block_len}}, + kv_head_{kv_head}, + block_len_divmod_{block_len_divmod}, + cache_block_offset_{cache_block_offset} + { + } + + template + CUTE_DEVICE auto row(int token) const + { + int page_index; + int page_token; + block_len_divmod_(page_index, page_token, token); + char* page = pages_[page_index]; + const auto byte_offset = Value ? layout_.v_data(kv_head_, page_token) : + layout_.k_data(kv_head_, page_token); + const T* row_ptr = reinterpret_cast(page + cache_block_offset_ + byte_offset); + return cute::make_tensor(cute::make_gmem_ptr(row_ptr), cute::make_layout(cute::Int{})); + } + + template + CUTE_DEVICE auto page(int page_index) const + { + char* page = pages_[page_index]; + const auto byte_offset = Value ? layout_.v_data(kv_head_, 0) : + layout_.k_data(kv_head_, 0); + const T* page_ptr = reinterpret_cast(page + cache_block_offset_ + byte_offset); + return cute::make_tensor( + cute::make_gmem_ptr(page_ptr), + cute::make_layout( + cute::make_shape(cute::Int{}, cute::Int{}), + cute::make_stride(cute::Int{}, cute::_1{}))); + } + +private: + char* const* pages_; + Layout layout_; + int kv_head_; + cutlass::FastDivmod block_len_divmod_; + int cache_block_offset_; +}; + +template +CUTE_DEVICE void copy_paged_tile(const PagedKv& cache, + int key_begin, + int key_limit, + SmemTensor destination) +{ + using Copy = KvTileCopy; + auto logical_tile = cute::make_identity_tensor(typename Policy::KvShape{}); + auto tiled_copy = typename Copy::TiledCopy{}; + auto thread_copy = tiled_copy.get_thread_slice(threadIdx.x); + auto coordinates = cute::group_modes< + 1, cute::rank_v>( + thread_copy.partition_S(logical_tile)); + auto destinations = cute::group_modes< + 1, cute::rank_v>( + thread_copy.partition_D(destination)); + + const int d_begin = cute::get<1>(coordinates(cute::_0{}, cute::_0{})); + auto copy_atom = typename Copy::Atom{}; + CUTE_UNROLL + for (int access = 0; access < Copy::AccessCount; ++access) { + const int key = key_begin + cute::get<0>(coordinates(cute::_0{}, access)); + const bool valid = key < key_limit; + const int source_key = valid ? key : key_begin; + auto source_row = cache.template row(source_key); + auto source = cute::make_tensor( + cute::make_gmem_ptr(&source_row(d_begin)), + cute::make_layout(cute::Int{})); + cute::copy(copy_atom.with(valid), source, destinations(cute::_, access)); + } +} + +template +CUTE_DEVICE void copy_paged_page_impl(const PagedKv& cache, + int page_index, + int key_begin, + int key_limit, + SmemTensor destination, + int thread_idx) +{ + using Copy = KvTileCopy; + auto logical_tile = cute::make_identity_tensor(typename Policy::KvShape{}); + auto source_page = cache.template page(page_index); + auto tiled_copy = typename Copy::TiledCopy{}; + auto thread_copy = tiled_copy.get_thread_slice(thread_idx); + auto coordinates = cute::group_modes< + 1, cute::rank_v>( + thread_copy.partition_S(logical_tile)); + auto sources = cute::group_modes< + 1, cute::rank_v>( + thread_copy.partition_S(source_page)); + auto destinations = cute::group_modes< + 1, cute::rank_v>( + thread_copy.partition_D(destination)); + + auto copy_atom = typename Copy::Atom{}; + CUTE_UNROLL + for (int access = 0; access < Copy::AccessCount; ++access) { + if constexpr (GuardTail) { + const int key = key_begin + cute::get<0>(coordinates(cute::_0{}, access)); + cute::copy(copy_atom.with(key < key_limit), + sources(cute::_, access), + destinations(cute::_, access)); + } + else { + cute::copy(copy_atom, sources(cute::_, access), destinations(cute::_, access)); + } + } +} + +template +CUTE_DEVICE void copy_paged_page(const PagedKv& cache, + int page_index, + int key_begin, + int key_limit, + SmemTensor destination, + int thread_idx = threadIdx.x) +{ + copy_paged_page_impl( + cache, page_index, key_begin, key_limit, destination, thread_idx); +} + +template +CUTE_DEVICE void copy_paged_page_full(const PagedKv& cache, + int page_index, + SmemTensor destination, + int thread_idx = threadIdx.x) +{ + copy_paged_page_impl( + cache, page_index, 0, 0, destination, thread_idx); +} + +} // namespace turbomind::verification_attention diff --git a/src/turbomind/kernels/attention/verification/policy_sm90.cuh b/src/turbomind/kernels/attention/verification/policy_sm90.cuh new file mode 100644 index 0000000000..0eeb71f65b --- /dev/null +++ b/src/turbomind/kernels/attention/verification/policy_sm90.cuh @@ -0,0 +1,159 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#pragma once + +#include + +#include + +#include +#include + +namespace turbomind::verification_attention { + +template +struct Sm90PolicyBase { + using MmaTraits = cute::MMA_Traits; + + static constexpr int HeadDim = HeadDim_; + static constexpr int MTile = MTile_; + static constexpr int MinThreads = 128; + static constexpr int MmaM = cute::size<0>(typename MmaTraits::Shape_MNK{}); + static constexpr int MmaThreads = cute::size(typename MmaTraits::ThrID{}); + static constexpr int MmaAtomM = MTile / MmaM; + static constexpr int MmaMThreads = MmaThreads * MmaAtomM; + static constexpr int MmaAtomN = (MinThreads + MmaMThreads - 1) / MmaMThreads; + static constexpr int Threads = MmaMThreads * MmaAtomN; + static constexpr int Stages = 3; + static constexpr int PvNtile = 32; + + static_assert(MTile % MmaM == 0); + + using QShape = cute::Shape, cute::Int>; + using OutputShape = QShape; +}; + +template +struct Sm90Policy: Sm90PolicyBase { + using Base = Sm90PolicyBase; + static constexpr int HeadDim = Base::HeadDim; + static constexpr int MTile = Base::MTile; + static constexpr int PvNtile = Base::PvNtile; + static constexpr int KeyTile = HeadDim == 128 ? 64 : 32; + static constexpr int PvTileCount = HeadDim / PvNtile; + + using KvShape = cute::Shape, cute::Int>; + using ScoreShape = cute::Shape, cute::Int>; + using PvShape = cute::Shape, cute::Int>; +}; + +template +struct VectorTileCopy { + using Atom = Atom_; + using T = typename Atom::ValType; + + static constexpr int ValuesPerAccess = 16 / sizeof(T); + static constexpr int VectorsPerRow = Columns / ValuesPerAccess; + static constexpr int RowsPerThreadTile = Policy::Threads / VectorsPerRow; + static constexpr int AccessCount = Rows / RowsPerThreadTile; + + using ThreadLayout = cute::Layout< + cute::Shape, cute::Int>, + cute::Stride, cute::_1>>; + using ValueLayout = cute::Layout>>; + using TiledCopy = decltype(cute::make_tiled_copy(Atom{}, ThreadLayout{}, ValueLayout{})); + + static_assert(Columns % ValuesPerAccess == 0); + static_assert(Policy::Threads % VectorsPerRow == 0); + static_assert(Rows % RowsPerThreadTile == 0); + static_assert(cute::size(ThreadLayout{}) == Policy::Threads); +}; + +template +using QTileCopy = VectorTileCopy, T>, + Policy::MTile, + Policy::HeadDim, + Policy>; + +template +using KvTileCopy = VectorTileCopy< + cute::Copy_Atom, T>, + Policy::KeyTile, + Policy::HeadDim, + Policy>; + +template +using OutputTileCopy = QTileCopy; + +template +struct MmaOperation; + +template<> +struct MmaOperation { + using Type = cute::SM80_16x8x16_F32F16F16F32_TN; +}; + +template<> +struct MmaOperation { + using Type = cute::SM80_16x8x16_F32BF16BF16F32_TN; +}; + +template +struct Sm90Mma { + using AtomLayout = cute::Layout< + cute::Shape, cute::Int, cute::_1>>; + using MmaOp = typename MmaOperation::Type; + + using QK = decltype(cute::make_tiled_mma( + MmaOp{}, + AtomLayout{}, + cute::Tile, cute::Int, cute::_16>{})); + using PV = decltype(cute::make_tiled_mma( + MmaOp{}, + AtomLayout{}, + cute::Tile, cute::Int, cute::_16>{})); + + using CopyQkA = decltype(cute::make_tiled_copy_A( + cute::Copy_Atom{}, QK{})); + using CopyQkB = decltype(cute::make_tiled_copy_B( + cute::Copy_Atom{}, QK{})); + using CopyPvA = decltype(cute::make_tiled_copy_A( + cute::Copy_Atom{}, PV{})); + using CopyPvBAtom = cute::Copy_Atom; + using CopyPvB = decltype(cute::make_tiled_copy_B(CopyPvBAtom{}, PV{})); + using StoreProbability = decltype(cute::make_tiled_copy_C( + cute::Copy_Atom{}, QK{})); +}; + +using SmemAtom = decltype(cute::composition( + cute::Swizzle<3, 3, 3>{}, + cute::Layout>, + cute::Stride>>{})); + +using SmemAtomNarrow = decltype(cute::composition( + cute::Swizzle<3, 3, 3>{}, + cute::Layout>, + cute::Stride>>{})); + +template +using SmemAtomFor = std::conditional_t; + +template +using SmemLayout2D = decltype(cute::tile_to_shape( + SmemAtomFor{}, cute::Shape, cute::Int>{})); + +template +using SmemLayout3D = decltype(cute::tile_to_shape( + SmemAtom{}, cute::Shape, cute::Int, cute::Int>{})); + +template +union SharedStorage { + alignas(16) T q[Policy::MTile * Policy::HeadDim]; + struct { + alignas(16) T k[Policy::Stages * Policy::KeyTile * Policy::HeadDim]; + alignas(16) T v[Policy::Stages * Policy::KeyTile * Policy::HeadDim]; + alignas(16) T probability[Policy::MTile * Policy::KeyTile]; + } body; +}; + +} // namespace turbomind::verification_attention diff --git a/src/turbomind/kernels/attention/verification/python_bind.cpp b/src/turbomind/kernels/attention/verification/python_bind.cpp new file mode 100644 index 0000000000..32367c375c --- /dev/null +++ b/src/turbomind/kernels/attention/verification/python_bind.cpp @@ -0,0 +1,365 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#include +#include + +#include + +#include + +#include "src/turbomind/kernels/attention/kv_cache_utils_v2.h" +#include "src/turbomind/kernels/attention/verification/attention.h" +#include "src/turbomind/python/attention_component_bindings.h" +#include "src/turbomind/python/eagle3_dlpack_internal.h" +#include "src/turbomind/utils/cuda_utils.h" + +namespace py = pybind11; + +namespace turbomind::python { +namespace { + +int GetCudaOrdinal(py::handle tensor) +{ + return tensor.attr("__dlpack_device__")().cast()[1].cast(); +} + +template +int LaunchVerificationAttention(py::handle prefix_k_object, + py::handle prefix_v_object, + py::handle prefix_offsets_object, + int max_history_length, + py::handle packed_qkv_object, + py::handle q_bias_object, + py::handle output_object, + py::handle cache_storage_object, + py::handle block_ptrs_object, + py::handle block_ptr_offsets_object, + py::handle q_offsets_object, + py::handle k_offsets_object, + py::handle finished_object, + py::handle partial_o_object, + py::handle partial_ml_object, + int query_head_count, + int kv_head_count, + int head_dim, + int block_len, + int max_query_length, + int max_key_length, + int window_size, + int requested_max_split_count, + int rope_type, + int rope_dim, + float rope_base, + float rope_factor, + int mrope_mode, + int mrope_section_t, + int mrope_section_h, + int mrope_section_w, + py::handle mrope_position_ids_object, + py::handle mrope_position_delta_object, + py::handle mrope_length_object, + uintptr_t stream_ptr) +{ + auto stream = reinterpret_cast(stream_ptr); + + auto prefix_k = detail::ConsumeDLPackWithStrides(prefix_k_object, stream_ptr); + auto prefix_v = detail::ConsumeDLPackWithStrides(prefix_v_object, stream_ptr); + auto prefix_offsets = detail::ConsumeDLPackWithStrides(prefix_offsets_object, stream_ptr); + auto packed_qkv = detail::ConsumeDLPackWithStrides(packed_qkv_object, stream_ptr); + auto q_bias = detail::ConsumeDLPackWithStrides(q_bias_object, stream_ptr); + auto output = detail::ConsumeDLPackWithStrides(output_object, stream_ptr); + auto cache_storage = + detail::ConsumeDLPackWithStrides(cache_storage_object, stream_ptr); + (void)cache_storage; + auto block_ptrs = detail::ConsumeDLPackWithStrides(block_ptrs_object, stream_ptr); + auto block_ptr_offsets = detail::ConsumeDLPackWithStrides(block_ptr_offsets_object, stream_ptr); + auto q_offsets = detail::ConsumeDLPackWithStrides(q_offsets_object, stream_ptr); + auto k_offsets = detail::ConsumeDLPackWithStrides(k_offsets_object, stream_ptr); + auto finished = detail::ConsumeDLPackWithStrides(finished_object, stream_ptr); + auto partial_o = detail::ConsumeDLPackWithStrides(partial_o_object, stream_ptr); + auto partial_ml = detail::ConsumeDLPackWithStrides(partial_ml_object, stream_ptr); + auto mrope_position_ids = detail::ConsumeDLPackWithStrides(mrope_position_ids_object, stream_ptr); + auto mrope_position_delta = + detail::ConsumeDLPackWithStrides(mrope_position_delta_object, stream_ptr); + auto mrope_length = detail::ConsumeDLPackWithStrides(mrope_length_object, stream_ptr); + + RopeKernelParam rope{}; + rope.type = static_cast(rope_type); + rope.dim = rope_dim; + rope.scale_factor = + rope.type == RopeType::kNull ? 0.f : -std::log2(rope_base) / rope_dim; + rope.inv_factor = rope_factor != 0.f ? 1.f / rope_factor : 1.f; + rope.mrope_mode = static_cast(mrope_mode); + if (rope.mrope_mode != MropeMode::kNone) { + rope.mrope.section = make_int3(mrope_section_t, mrope_section_h, mrope_section_w); + rope.mrope.stride = static_cast(mrope_position_ids.stride(0)); + rope.mrope.position_ids = mrope_position_ids.data(); + rope.mrope.position_delta = mrope_position_delta.data(); + rope.mrope.length = mrope_length.data(); + } + + const int batch_size = static_cast(q_offsets.shape(0) - 1); + const int64_t prefix_head_stride = prefix_k.stride(0) / head_dim; + invokeProcessKV_v2(reinterpret_cast(block_ptrs.data()), + prefix_k.data(), + prefix_v.data(), + nullptr, + nullptr, + prefix_offsets.data(), + prefix_offsets.data(), + block_ptr_offsets.data(), + nullptr, + nullptr, + rope, + 0, + prefix_head_stride, + 1, + prefix_head_stride, + block_len, + 0, + 0, + cutlass::FastDivmod(1), + max_history_length, + kv_head_count, + head_dim, + batch_size, + 0, + stream); + + const int64_t qkv_stride = packed_qkv.stride(0); + const int64_t qkv_head_stride = qkv_stride / head_dim; + const T* packed = packed_qkv.data(); + const T* submitted_k = packed + query_head_count * head_dim; + const T* submitted_v = submitted_k + kv_head_count * head_dim; + invokeProcessKV_v2(reinterpret_cast(block_ptrs.data()), + submitted_k, + submitted_v, + nullptr, + nullptr, + q_offsets.data(), + k_offsets.data(), + block_ptr_offsets.data(), + nullptr, + finished.data(), + rope, + 0, + qkv_head_stride, + 1, + qkv_head_stride, + block_len, + 0, + 0, + cutlass::FastDivmod(1), + max_query_length, + kv_head_count, + head_dim, + batch_size, + 0, + stream); + + verification_attention::Arguments arguments{}; + arguments.out = output.data(); + arguments.q = packed; + arguments.q_bias = q_bias.shape(0) ? q_bias.data() : nullptr; + arguments.q_stride = qkv_stride; + arguments.block_ptrs = + reinterpret_cast(block_ptrs.data()); + arguments.block_ptr_offsets = block_ptr_offsets.data(); + arguments.q_offsets = q_offsets.data(); + arguments.k_offsets = k_offsets.data(); + arguments.finished = finished.data(); + arguments.request_count = batch_size; + arguments.query_count = static_cast(packed_qkv.shape(0)); + arguments.query_offset = 0; + arguments.max_query_length = max_query_length; + arguments.max_key_length = max_key_length; + arguments.query_head_count = query_head_count; + arguments.kv_head_count = kv_head_count; + arguments.query_group_size = query_head_count / kv_head_count; + arguments.query_group_size_divmod = cutlass::FastDivmod(arguments.query_group_size); + arguments.head_dim = head_dim; + arguments.block_len = block_len; + arguments.block_len_divmod = cutlass::FastDivmod(block_len); + arguments.cache_block_offset = 0; + arguments.window_size = window_size ? window_size : (256 << 20); + arguments.qk_scale_log2 = + std::log2(std::exp(1.f)) / std::sqrt(static_cast(head_dim)); + arguments.rope = rope; + arguments.partial_o = partial_o.data(); + arguments.partial_ml = partial_ml.data(); + arguments.data_type = packed_qkv.dtype(); + arguments.stream = stream; + + const int m_slices = (max_query_length * arguments.query_group_size + + verification_attention::CtaM(arguments) - 1) + / verification_attention::CtaM(arguments); + const int base_cta_count = batch_size * kv_head_count * m_slices; + arguments.split_count = verification_attention::choose_split_count(arguments.query_count, + base_cta_count, + arguments.max_key_length, + verification_attention::KeyTile(arguments), + static_cast(partial_o.shape(0)), + requested_max_split_count, + getSMCount()); + verification_attention::run(arguments); + return arguments.split_count; +} + +} // namespace + +void BindVerificationAttention(py::module_& module) +{ +#if TM_BUILD_VERIFICATION_ATTENTION_SM90 + module.def( + "verification_attention", + [](py::handle prefix_k, + py::handle prefix_v, + py::handle prefix_offsets, + int max_history_length, + py::handle packed_qkv, + py::handle q_bias, + py::handle output, + py::handle cache_storage, + py::handle block_ptrs, + py::handle block_ptr_offsets, + py::handle q_offsets, + py::handle k_offsets, + py::handle finished, + py::handle partial_o, + py::handle partial_ml, + int query_head_count, + int kv_head_count, + int head_dim, + int block_len, + int max_query_length, + int max_key_length, + int window_size, + int requested_max_split_count, + int rope_type, + int rope_dim, + float rope_base, + float rope_factor, + int mrope_mode, + int mrope_section_t, + int mrope_section_h, + int mrope_section_w, + py::handle mrope_position_ids, + py::handle mrope_position_delta, + py::handle mrope_length, + uintptr_t stream_ptr) { + CudaDeviceGuard guard{GetCudaOrdinal(packed_qkv)}; + const auto dtype = detail::ConsumeDLPackWithStrides(packed_qkv, stream_ptr).dtype(); + if (dtype == DataType::kHalf) { + return LaunchVerificationAttention(prefix_k, + prefix_v, + prefix_offsets, + max_history_length, + packed_qkv, + q_bias, + output, + cache_storage, + block_ptrs, + block_ptr_offsets, + q_offsets, + k_offsets, + finished, + partial_o, + partial_ml, + query_head_count, + kv_head_count, + head_dim, + block_len, + max_query_length, + max_key_length, + window_size, + requested_max_split_count, + rope_type, + rope_dim, + rope_base, + rope_factor, + mrope_mode, + mrope_section_t, + mrope_section_h, + mrope_section_w, + mrope_position_ids, + mrope_position_delta, + mrope_length, + stream_ptr); + } + return LaunchVerificationAttention(prefix_k, + prefix_v, + prefix_offsets, + max_history_length, + packed_qkv, + q_bias, + output, + cache_storage, + block_ptrs, + block_ptr_offsets, + q_offsets, + k_offsets, + finished, + partial_o, + partial_ml, + query_head_count, + kv_head_count, + head_dim, + block_len, + max_query_length, + max_key_length, + window_size, + requested_max_split_count, + rope_type, + rope_dim, + rope_base, + rope_factor, + mrope_mode, + mrope_section_t, + mrope_section_h, + mrope_section_w, + mrope_position_ids, + mrope_position_delta, + mrope_length, + stream_ptr); + }, + py::arg("prefix_k"), + py::arg("prefix_v"), + py::arg("prefix_offsets"), + py::arg("max_history_length"), + py::arg("packed_qkv"), + py::arg("q_bias"), + py::arg("output"), + py::arg("cache_storage"), + py::arg("block_ptrs"), + py::arg("block_ptr_offsets"), + py::arg("q_offsets"), + py::arg("k_offsets"), + py::arg("finished"), + py::arg("partial_o"), + py::arg("partial_ml"), + py::arg("query_head_count"), + py::arg("kv_head_count"), + py::arg("head_dim"), + py::arg("block_len"), + py::arg("max_query_length"), + py::arg("max_key_length"), + py::arg("window_size"), + py::arg("requested_max_split_count"), + py::arg("rope_type"), + py::arg("rope_dim"), + py::arg("rope_base"), + py::arg("rope_factor"), + py::arg("mrope_mode"), + py::arg("mrope_section_t"), + py::arg("mrope_section_h"), + py::arg("mrope_section_w"), + py::arg("mrope_position_ids"), + py::arg("mrope_position_delta"), + py::arg("mrope_length"), + py::arg("stream_ptr")); +#else + (void)module; +#endif +} + +} // namespace turbomind::python diff --git a/src/turbomind/kernels/attention/verification/reduce.cu b/src/turbomind/kernels/attention/verification/reduce.cu new file mode 100644 index 0000000000..0d9fcf8276 --- /dev/null +++ b/src/turbomind/kernels/attention/verification/reduce.cu @@ -0,0 +1,136 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#include "src/turbomind/kernels/attention/verification/kernel_sm80.cuh" + +#include + +namespace turbomind::verification_attention { + +template +__global__ void ReduceSplits(T* out_ptr, + const float* partial_o_ptr, + const float* partial_ml_ptr, + int query_count, + int query_offset, + int query_head_count, + int split_count) +{ + __shared__ float maxima[128]; + __shared__ float sums[128]; + __shared__ float weights[128]; + __shared__ float global_max; + __shared__ float global_sum; + + using BlockReduce = cub::BlockReduce; + union ReduceStorage { + typename BlockReduce::TempStorage maximum; + typename BlockReduce::TempStorage sum; + }; + __shared__ ReduceStorage reduce_storage; + + auto partial_o = cute::make_tensor( + cute::make_gmem_ptr(partial_o_ptr), + cute::make_layout( + cute::make_shape(query_count, split_count, query_head_count, cute::Int{}), + cute::make_stride(split_count * query_head_count * HeadDim, + query_head_count * HeadDim, + cute::Int{}, + cute::_1{}))); + auto partial_ml = cute::make_tensor( + cute::make_gmem_ptr(partial_ml_ptr), + cute::make_layout( + cute::make_shape(query_count, split_count, query_head_count, cute::_2{}), + cute::make_stride(split_count * query_head_count * 2, + query_head_count * 2, + cute::_2{}, + cute::_1{}))); + auto out = cute::make_tensor( + cute::make_gmem_ptr(out_ptr), + cute::make_layout( + cute::make_shape(query_count, query_head_count, cute::Int{}), + cute::make_stride(query_head_count * HeadDim, cute::Int{}, cute::_1{}))); + + using ThreadLayout = cute::Layout>; + auto thread_coordinates = cute::make_identity_tensor(cute::Shape{}); + auto owned_coordinate = cute::local_partition(thread_coordinates, ThreadLayout{}, threadIdx.x); + static_assert(cute::size(owned_coordinate) == 1); + + const int local_query = blockIdx.x; + const int head = blockIdx.y; + const int coordinate = cute::get<0>(owned_coordinate(cute::_0{})); + const int split = coordinate; + + float thread_max = -CUDART_INF_F; + if (split < split_count) { + maxima[split] = partial_ml(local_query, split, head, cute::_0{}); + sums[split] = partial_ml(local_query, split, head, cute::_1{}); + thread_max = maxima[split]; + } + thread_max = BlockReduce(reduce_storage.maximum).Reduce(thread_max, cub::Max{}); + if (threadIdx.x == 0) { + global_max = thread_max; + } + __syncthreads(); + + float thread_sum = 0.f; + if (split < split_count) { + const float weight = maxima[split] == -CUDART_INF_F + ? 0.f + : exp2f(maxima[split] - global_max); + weights[split] = weight; + thread_sum = weight * sums[split]; + } + thread_sum = BlockReduce(reduce_storage.sum).Sum(thread_sum); + if (threadIdx.x == 0) { + global_sum = thread_sum; + } + __syncthreads(); + + const int d = coordinate; + if (d < HeadDim) { + float numerator = 0.f; + for (int split_index = 0; split_index < split_count; ++split_index) { + numerator += weights[split_index] * partial_o(local_query, split_index, head, d); + } + out(local_query + query_offset, head, d) = + global_sum == 0.f ? T(0) : static_cast(numerator / global_sum); + } +} + +template +void DispatchReduce(const Arguments& arguments) +{ + dim3 grid(arguments.query_count, arguments.query_head_count); + if (arguments.head_dim == 128) { + ReduceSplits<<>>( + static_cast(arguments.out), + arguments.partial_o, + arguments.partial_ml, + arguments.query_count, + arguments.query_offset, + arguments.query_head_count, + arguments.split_count); + } + else { + ReduceSplits<<>>( + static_cast(arguments.out), + arguments.partial_o, + arguments.partial_ml, + arguments.query_count, + arguments.query_offset, + arguments.query_head_count, + arguments.split_count); + } +} + +void Reduce(const Arguments& arguments) +{ + if (arguments.data_type == DataType::kHalf) { + DispatchReduce(arguments); + } + else { + DispatchReduce(arguments); + } +} + +} // namespace turbomind::verification_attention diff --git a/src/turbomind/kernels/attention/verification/stub.cc b/src/turbomind/kernels/attention/verification/stub.cc new file mode 100644 index 0000000000..d22cc3bfa5 --- /dev/null +++ b/src/turbomind/kernels/attention/verification/stub.cc @@ -0,0 +1,21 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#include "src/turbomind/kernels/attention/verification/attention.h" + +namespace turbomind::verification_attention { + +bool supports(const Capability&) +{ + return false; +} + +int choose_split_count(int, int, int, int, int, int, int) +{ + return 1; +} + +void run(const Arguments&) +{ +} + +} // namespace turbomind::verification_attention diff --git a/src/turbomind/kernels/ban_bad_words.cu b/src/turbomind/kernels/ban_bad_words.cu index 7f1bb6b452..19a2774c5e 100644 --- a/src/turbomind/kernels/ban_bad_words.cu +++ b/src/turbomind/kernels/ban_bad_words.cu @@ -54,12 +54,17 @@ __global__ void BanBadWordsKernel(T* logits, const int* const* token_ids_ptrs, const int* sequence_length, const int* bad_words, + const bool* logits_active, size_t bad_words_len, int vocab_size) { const int id = blockIdx.x * blockDim.x + threadIdx.x; const int batch_idx = blockIdx.y; + if (logits_active != nullptr && !logits_active[batch_idx]) { + return; + } + const int* base_bad_words = bad_words + batch_idx * 2 * bad_words_len; const int* base_bad_words_offsets = base_bad_words + bad_words_len; @@ -101,6 +106,7 @@ void BanBadWords(Tensor& logits, const Buffer_ token_ids_ptrs, const Buffer_& sequence_length, const Tensor_& bad_words, + const bool* logits_active, cudaStream_t stream) { @@ -117,6 +123,7 @@ void BanBadWords(Tensor& logits, token_ids_ptrs.data(), sequence_length.data(), bad_words.data(), + logits_active, bad_words_len, vocab_size); }; diff --git a/src/turbomind/kernels/ban_bad_words.h b/src/turbomind/kernels/ban_bad_words.h index eb2c1e353d..297cf3afd1 100644 --- a/src/turbomind/kernels/ban_bad_words.h +++ b/src/turbomind/kernels/ban_bad_words.h @@ -26,6 +26,7 @@ void BanBadWords(Tensor& logits, const Buffer_ token_ids_ptrs, const Buffer_& sequence_length, const Tensor_& bad_words, + const bool* logits_active, cudaStream_t stream); } // namespace turbomind diff --git a/src/turbomind/kernels/copy/copy.cc b/src/turbomind/kernels/copy/copy.cc index e78fe98764..21a05c06d0 100644 --- a/src/turbomind/kernels/copy/copy.cc +++ b/src/turbomind/kernels/copy/copy.cc @@ -3,6 +3,7 @@ #include "src/turbomind/core/logger.h" #include +#include #include #include @@ -15,56 +16,32 @@ void VectorizedCopy( void TransposeCopy( const void* data_a, void* data_b, const Layout& a, const Layout& b, DataType dtype, cudaStream_t stream); -// Merge adjacent batch dims (positions ≥ 2) of (a, b) when their strides are -// proportional in BOTH a and b. Single forward pass over positions 3..rank-1. -// -// Precondition: a and b have the same shape and rank, with positions 0 and 1 -// being the (I, J) transpose pair (not coalesceable). Only positions 2.. are -// considered batch dims. -// -// Why the joint-proportionality requirement? -// A batch dim is shared between src and dst: the kernel decodes blockIdx.z -// once into a multi-dim batch coord and dots it with the batch strides on -// BOTH sides to compute the per-block src and dst pointer offsets. Merging -// two adjacent batch dims into one collapses that pair of coords into a -// single linear index whose decode (idx / inner_shape, idx % inner_shape) -// only reproduces the original (outer_idx, inner_idx) — and therefore the -// original outer_idx*outer_stride + inner_idx*inner_stride offset — when -// outer_stride == inner_shape * inner_stride. Since the same merged index -// is dotted into both layouts, that proportionality must hold in BOTH a -// and b; otherwise the merged single-dim decode would land at different -// positions in src vs dst and produce wrong results. -static std::pair coalesce_batch_dims(const Layout& a, const Layout& b) +// Both kernels decode one coordinate for source and destination. Remove +// singleton dimensions and merge contiguous dimensions jointly, in the +// innermost-first order used by the kernels. The first `begin` dimensions +// are preserved when coalescing only the batch axes of a transpose. +static std::pair coalesce_copy_dims(const Layout& a, const Layout& b, int begin = 0) { - const int rank = a.rank(); - if (rank < 4) - return {a, b}; // need ≥ 2 batch dims to merge - - std::vector ash(a.shape().begin(), a.shape().begin() + 3); - std::vector ast(a.stride().begin(), a.stride().begin() + 3); - std::vector bsh(b.shape().begin(), b.shape().begin() + 3); - std::vector bst(b.stride().begin(), b.stride().begin() + 3); - - for (int i = 3; i < rank; ++i) { - const ssize_t ai_sh = a.shape(i), ai_st = a.stride(i); - const ssize_t bi_sh = b.shape(i), bi_st = b.stride(i); - - // Merge with the previously accumulated batch dim if its stride equals - // shape * stride of that dim, in BOTH a and b. - if (ai_st == ash.back() * ast.back() && bi_st == bsh.back() * bst.back()) { - ash.back() *= ai_sh; - bsh.back() *= bi_sh; - // strides at the back stay unchanged (they remain the inner stride) + std::vector shape, src_stride, dst_stride; + for (int i = 0; i < a.rank(); ++i) { + if (i >= begin && a.shape(i) == 1) { + continue; + } + if (shape.size() > static_cast(begin) + && a.stride(i) == shape.back() * src_stride.back() + && b.stride(i) == shape.back() * dst_stride.back()) { + shape.back() *= a.shape(i); } else { - ash.push_back(ai_sh); - ast.push_back(ai_st); - bsh.push_back(bi_sh); - bst.push_back(bi_st); + shape.push_back(a.shape(i)); + src_stride.push_back(a.stride(i)); + dst_stride.push_back(b.stride(i)); } } - - return {Layout{ash, ast}, Layout{bsh, bst}}; + if (shape.empty()) { + return {Layout{{1}, {1}}, Layout{{1}, {1}}}; + } + return {Layout{shape, src_stride}, Layout{shape, dst_stride}}; } // ============================================================================ @@ -75,33 +52,34 @@ void GenericCopy(const Tensor& src, Tensor& dst, cudaStream_t stream) auto a = src.layout(); auto b = dst.layout(); - TM_CHECK_EQ(a.size(), b.size()) << "GenericCopy: src and dst must have the same number of elements"; + TM_CHECK(src.dtype() == dst.dtype()) << "GenericCopy: src and dst must have the same dtype"; + TM_CHECK(a.shape() == b.shape()) << "GenericCopy: src and dst must have the same shape"; + TM_CHECK_GT(byte_size(src.dtype()), 0) << "GenericCopy: sub-byte elements are unsupported"; + if (a.size() == 0) { + return; + } - // Sort strides ascending so innermost (fastest-varying) dim is first + // Put physical source axes first, keeping broadcast axes outside them. + // Apply exactly the same permutation and coalescing to both layouts. std::vector idxs(a.rank()); std::iota(idxs.begin(), idxs.end(), 0); - std::sort(idxs.begin(), idxs.end(), [&](int i, int j) { return a.stride()[i] < a.stride()[j]; }); + std::stable_sort(idxs.begin(), idxs.end(), [&](int i, int j) { + if ((a.stride(i) == 0) != (a.stride(j) == 0)) { + return a.stride(i) != 0; + } + return a.stride(i) < a.stride(j); + }); a = a.permute(idxs); b = b.permute(idxs); - a = a.coalesce(); - b = b.coalesce(); - - int rank = std::max(a.rank(), b.rank()); - - if (a.rank() < rank) { - a = a.view(b.shape()); - } - else if (b.rank() < rank) { - b = b.view(a.shape()); - } + std::tie(a, b) = coalesce_copy_dims(a, b); + const int rank = a.rank(); const DataType dtype = src.dtype(); // --- Transpose detection (2D + batched) --- - // After the src-stride-ascending sort above, position 0 holds the smallest - // src stride (innermost). We dispatch to TransposeCopy when: + // After joint normalization, we dispatch to TransposeCopy when: // - position 0 has src stride 1 (call it I), // - some position J ∈ [1, rank-1] has dst stride 1, // - both shape(0) and shape(J) are divisible by the per-dtype tile. @@ -120,23 +98,24 @@ void GenericCopy(const Tensor& src, Tensor& dst, cudaStream_t stream) bool is_transpose = (J >= 1) && (a.stride(0) == 1) && (a.stride(J) > 1) && (b.stride(0) > 1) && (a.shape(0) % kTileDim == 0) && (a.shape(J) % kTileDim == 0); + // TransposeCopy uses 16-byte atoms. Every row and batch base must satisfy + // that alignment; otherwise the generic path chooses a legal copy width. + is_transpose = is_transpose && reinterpret_cast(src.raw_data()) % 16 == 0 + && reinterpret_cast(dst.raw_data()) % 16 == 0; + for (int i = 0; is_transpose && i < rank; ++i) { + is_transpose = (i == 0 || byte_size(dtype, a.stride(i)) % 16 == 0) + && (i == J || byte_size(dtype, b.stride(i)) % 16 == 0); + } + if (is_transpose) { if (J != 1) { a = a.transpose(1, J); b = b.transpose(1, J); } - std::tie(a, b) = coalesce_batch_dims(a, b); - - // Compute total batch (product of post-coalesce batch dims, positions ≥ 2). - int64_t total_batch = 1; - for (int i = 2; i < a.rank(); ++i) - total_batch *= a.shape(i); - - // Dispatch only when the kernel can handle it: - // 1. post-coalesce rank ≤ 4 (only 2/3/4 are instantiated in TransposeCopy host), - // 2. total_batch ≤ gridDim.z hardware limit (65535 on all current archs). - // Otherwise, fall through to VectorizedCopy. - if (a.rank() <= 4 && total_batch <= 65535) { + std::tie(a, b) = coalesce_copy_dims(a, b, 2); + + // TransposeCopy maps all logical tile axes onto bounded linear launches. + if (a.rank() <= 4) { TransposeCopy(src.raw_data(), dst.raw_data(), a, b, dtype, stream); return; } diff --git a/src/turbomind/kernels/copy/copy.cu b/src/turbomind/kernels/copy/copy.cu index b7e2617aec..f219635888 100644 --- a/src/turbomind/kernels/copy/copy.cu +++ b/src/turbomind/kernels/copy/copy.cu @@ -20,16 +20,16 @@ using namespace cute; // from runtime shape/stride arrays. namespace detail { -template +template auto make_cute_shape_impl(const ssize_t* data, std::index_sequence) { - return make_shape(static_cast(data[Is])...); + return make_shape(static_cast(data[Is])...); } -template +template auto make_cute_shape(const ssize_t* data) { - return make_cute_shape_impl(data, std::make_index_sequence{}); + return make_cute_shape_impl(data, std::make_index_sequence{}); } template @@ -70,31 +70,25 @@ auto make_vec_factors() } // Compute thread partition: (T0, T1, ..., Tk-1) where T0*...*Tk-1 = 256. -// T0 is the largest power-of-2 <= shape[0]/kVec. -// Remaining threads are distributed across outer dims. +// Use power-of-two factors for inner dimensions and give the outermost +// dimension all remaining threads, so every thread belongs to this tile. template auto compute_thr_partition(const ssize_t* shape, int kVec) -> std::array { std::array partition{}; partition.fill(1); - // Inner dim: largest power-of-2 that divides 256 and <= shape[0]/kVec - int64_t max_inner = shape[0] / kVec; - ssize_t T0 = 256; - while (T0 > 1 && T0 > max_inner) { - T0 /= 2; - } - partition[0] = T0; - - // Distribute remaining threads across outer dims - ssize_t remaining = 256 / T0; - for (int i = 1; i < kRank; ++i) { - partition[i] = std::min(shape[i], remaining); - remaining /= partition[i]; - if (remaining < 1) { - remaining = 1; + ssize_t remaining = 256; + for (int i = 0; i < kRank - 1; ++i) { + const ssize_t extent = shape[i] / (i == 0 ? kVec : 1); + ssize_t threads = remaining; + while (threads > 1 && threads > extent) { + threads /= 2; } + partition[i] = threads; + remaining /= threads; } + partition[kRank - 1] = remaining; return partition; } @@ -169,10 +163,6 @@ void VectorizedCopy( auto align = [&](auto v) { alignment = std::gcd(alignment, v); }; - if (a.stride(0) > 1 || b.stride(0) > 1) { - alignment = byte_size(dtype); - } - align(byte_size(dtype, a.shape(0))); align(reinterpret_cast(data_a)); align(reinterpret_cast(data_b)); @@ -184,7 +174,12 @@ void VectorizedCopy( // --- vec_size computation --- const int elem_size = byte_size(dtype); - int vec_size = static_cast(alignment / std::max(1, elem_size)); + // A vector atom copies adjacent elements. Alignment alone cannot make a + // broadcast or strided axis contiguous, so those layouts use one element. + int vec_size = 1; + if (a.stride(0) == 1 && b.stride(0) == 1) { + vec_size = static_cast(alignment / elem_size); + } if (vec_size * elem_size > 16) { vec_size = 16 / elem_size; @@ -213,7 +208,9 @@ void VectorizedCopy( auto dst_strides = detail::make_cute_stride(b.stride().data()); auto partition_arr = detail::compute_thr_partition(a.shape().data(), kVec); - auto thr_partition = detail::make_cute_shape(partition_arr.data()); + // Thread coordinates are bounded by the 256-thread block; + // logical data shapes can exceed INT32_MAX after coalescing. + auto thr_partition = detail::make_cute_shape(partition_arr.data()); auto vec_factors = detail::make_vec_factors(); auto tile_sizes = transform(thr_partition, vec_factors, [](auto tp, auto vf) { return tp * vf; }); diff --git a/src/turbomind/kernels/copy/copy.h b/src/turbomind/kernels/copy/copy.h index 2921a92d9f..c54cf5d9aa 100644 --- a/src/turbomind/kernels/copy/copy.h +++ b/src/turbomind/kernels/copy/copy.h @@ -4,6 +4,9 @@ namespace turbomind::core { +// Elementwise device copy between tensors of the same shape and dtype. +// Source broadcast strides are supported. Destination elements and the two +// buffers must not overlap. At most four jointly coalesced axes are supported. void GenericCopy(const Tensor& src, Tensor& dst, cudaStream_t stream); } // namespace turbomind::core diff --git a/src/turbomind/kernels/copy/transpose.cu b/src/turbomind/kernels/copy/transpose.cu index ec8600b362..129bf646a7 100644 --- a/src/turbomind/kernels/copy/transpose.cu +++ b/src/turbomind/kernels/copy/transpose.cu @@ -3,6 +3,8 @@ #include "src/turbomind/kernels/copy/copy.h" #include #include +#include +#include namespace turbomind::core { @@ -15,9 +17,13 @@ namespace kernel { extern __shared__ char smem_buf[]; -template +template __global__ void __launch_bounds__(256) - TransposeCopyKernel(cute::Tensor src, cute::Tensor dst) + TransposeCopyKernel(cute::Tensor src, + cute::Tensor dst, + TileGridShape tile_grid_shape, + int64_t tile_base) { using T = typename SrcEngine::value_type; static_assert(std::is_same_v, @@ -40,7 +46,14 @@ __global__ void __launch_bounds__(256) make_tensor(make_smem_ptr(smem_base + kTileDim * kStride), make_layout(make_shape(Int{}, Int{}), make_stride(Int{}, Int<1>{}))); - // Decode blockIdx.z → multi-dim batch coord → per-block pointer offsets. + // Preserve the original column, row, batch tile order in a linear grid. + // Logical row and batch counts no longer consume grid.y or grid.z. + const auto tile_coord = idx2crd(tile_base + int64_t(blockIdx.x), tile_grid_shape); + const auto tile_n = get<0>(tile_coord); + const auto tile_m = get<1>(tile_coord); + const auto batch = get<2>(tile_coord); + + // Decode the logical batch index into per-block pointer offsets. // The `if constexpr (kRank > 2)` guard is REQUIRED: cute::crd2idx / // cute::idx2crd are implemented with unary fold expressions of the form // `(... + crd2idx_inner(...))` over the shape's tuple_seq. For rank == 2 @@ -54,7 +67,7 @@ __global__ void __launch_bounds__(256) auto batch_shape = take<2, kRank>(shape(src)); auto src_batch_str = take<2, kRank>(stride(src)); auto dst_batch_str = take<2, kRank>(stride(dst)); - auto batch_coord = idx2crd(int64_t(blockIdx.z), batch_shape); + auto batch_coord = idx2crd(batch, batch_shape); src_off = crd2idx(batch_coord, batch_shape, src_batch_str); dst_off = crd2idx(batch_coord, batch_shape, dst_batch_str); } @@ -76,12 +89,8 @@ __global__ void __launch_bounds__(256) auto src_tiled = tiled_divide(src_2d, tiler); auto dst_tiled = tiled_divide(dst_2d, tiler); - // Bounds check on tile grid - if (blockIdx.y >= size<1>(src_tiled) || blockIdx.x >= size<2>(src_tiled)) - return; - - auto src_tile = src_tiled(make_coord(_, _), blockIdx.y, blockIdx.x); - auto dst_tile = dst_tiled(make_coord(_, _), blockIdx.y, blockIdx.x); + auto src_tile = src_tiled(make_coord(_, _), tile_m, tile_n); + auto dst_tile = dst_tiled(make_coord(_, _), tile_m, tile_n); // Phase 1: gmem(src) -> smem1, vectorize along dim 0 auto tc1 = make_tiled_copy(Copy_Atom, T>{}, @@ -118,7 +127,7 @@ namespace detail { template auto make_cute_shape_impl(const ssize_t* data, std::index_sequence) { - return make_shape(static_cast(data[Is])...); + return make_shape(static_cast(data[Is])...); } template @@ -161,9 +170,9 @@ auto make_unit_stride(const ssize_t* data) void TransposeCopy( const void* data_a, void* data_b, const Layout& a, const Layout& b, DataType dtype, cudaStream_t stream) { - const int rank = a.rank(); - int32_t M = static_cast(a.shape(0)); - int32_t N = static_cast(a.shape(1)); + const int rank = a.rank(); + const int64_t M = a.shape(0); + const int64_t N = a.shape(1); auto launch = [&](auto t, auto kvec, auto ktiledim, auto rank_c) { using T = decltype(t); @@ -189,11 +198,16 @@ void TransposeCopy( total_batch *= a.shape(i); constexpr int smem_bytes = 2 * kTileDim * (kTileDim + kVec) * sizeof(T); - dim3 grid(static_cast(N / kTileDim), - static_cast(M / kTileDim), - static_cast(total_batch)); - - kernel::TransposeCopyKernel<<>>(src_gmem, dst_gmem); + const auto tile_grid_shape = make_shape(N / kTileDim, M / kTileDim, total_batch); + const int64_t total_tiles = product(tile_grid_shape); + constexpr int64_t max_blocks = std::numeric_limits::max(); + // Keep every launch within grid.x's limit while retaining 64-bit + // logical tile indices. Normal workloads use a single launch. + for (int64_t tile_base = 0; tile_base < total_tiles; tile_base += max_blocks) { + const auto blocks = static_cast(std::min(total_tiles - tile_base, max_blocks)); + kernel::TransposeCopyKernel<<>>( + src_gmem, dst_gmem, tile_grid_shape, tile_base); + } }; auto dispatch_rank = [&](auto t, auto kvec, auto ktiledim) { diff --git a/src/turbomind/kernels/draft_carry_kernels.cu b/src/turbomind/kernels/draft_carry_kernels.cu new file mode 100644 index 0000000000..4f1500a8d5 --- /dev/null +++ b/src/turbomind/kernels/draft_carry_kernels.cu @@ -0,0 +1,81 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#include +#include + +#include + +#include "src/turbomind/kernels/draft_carry_kernels.h" + +namespace turbomind { +namespace { + +__global__ void SelectDraftCarryKernel(const unsigned char* local_residual, + const int* selected_local_rows, + const bool* candidate_active, + unsigned char* carry, + int local_token_num, + int row_bytes, + int first, + int last) +{ + const int candidate = blockIdx.x; + + auto* destination = reinterpret_cast(carry + static_cast(candidate) * row_bytes); + const int vector_count = row_bytes / sizeof(uint4); + + bool owned = candidate_active[candidate]; + int row = 0; + if (owned) { + row = selected_local_rows[candidate]; + owned = 0 <= row && row < local_token_num && first <= row && row < last; + } + + if (owned) { + const auto* source = reinterpret_cast(local_residual + static_cast(row) * row_bytes); + for (int i = threadIdx.x; i < vector_count; i += blockDim.x) { + destination[i] = source[i]; + } + } + else { + const uint4 zero{}; + for (int i = threadIdx.x; i < vector_count; i += blockDim.x) { + destination[i] = zero; + } + } +} + +} // namespace + +void invokeSelectDraftCarry(const void* local_residual, + const int* selected_local_rows, + const bool* candidate_active, + void* carry, + int local_token_num, + int candidate_count, + int hidden_size, + int element_bits, + int first, + int last, + cudaStream_t stream) +{ + if (candidate_count == 0) { + return; + } + + const int64_t row_bits = static_cast(hidden_size) * element_bits; + const int row_bytes = row_bits / 8; + + constexpr int block_size = 256; + SelectDraftCarryKernel<<>>( + static_cast(local_residual), + selected_local_rows, + candidate_active, + static_cast(carry), + local_token_num, + row_bytes, + first, + last); +} + +} // namespace turbomind diff --git a/src/turbomind/kernels/draft_carry_kernels.h b/src/turbomind/kernels/draft_carry_kernels.h new file mode 100644 index 0000000000..daa826635a --- /dev/null +++ b/src/turbomind/kernels/draft_carry_kernels.h @@ -0,0 +1,21 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#pragma once + +#include + +namespace turbomind { + +void invokeSelectDraftCarry(const void* local_residual, + const int* selected_local_rows, + const bool* candidate_active, + void* carry, + int local_token_num, + int candidate_count, + int hidden_size, + int element_bits, + int first, + int last, + cudaStream_t stream); + +} // namespace turbomind diff --git a/src/turbomind/kernels/draft_carry_python_bind.cpp b/src/turbomind/kernels/draft_carry_python_bind.cpp new file mode 100644 index 0000000000..6a7b01d462 --- /dev/null +++ b/src/turbomind/kernels/draft_carry_python_bind.cpp @@ -0,0 +1,64 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#include + +#include + +#include + +#include "src/turbomind/kernels/draft_carry_kernels.h" +#include "src/turbomind/python/eagle3_component_bindings.h" +#include "src/turbomind/python/eagle3_dlpack_internal.h" +#include "src/turbomind/utils/cuda_utils.h" + +namespace py = pybind11; + +namespace turbomind::python { +namespace { + +int GetCudaOrdinal(py::handle tensor) +{ + return tensor.attr("__dlpack_device__")().cast()[1].cast(); +} + +} // namespace + +void BindDraftCarry(py::module_& module) +{ + module.def( + "select_draft_carry", + [](py::handle local_residual_object, + py::handle selected_local_rows_object, + py::handle candidate_active_object, + py::handle carry_object, + int first, + int last, + uintptr_t stream_ptr) { + CudaDeviceGuard guard{GetCudaOrdinal(local_residual_object)}; + auto local_residual = detail::ConsumeDLPackWithStrides(local_residual_object, stream_ptr); + auto selected_local_rows = detail::ConsumeDLPackWithStrides(selected_local_rows_object, stream_ptr); + auto candidate_active = detail::ConsumeDLPackWithStrides(candidate_active_object, stream_ptr); + auto carry = detail::ConsumeDLPackWithStrides(carry_object, stream_ptr); + + invokeSelectDraftCarry(local_residual.data_or(static_cast(nullptr)), + selected_local_rows.data_or(static_cast(nullptr)), + candidate_active.data_or(static_cast(nullptr)), + carry.data_or(static_cast(nullptr)), + static_cast(local_residual.shape(0)), + static_cast(selected_local_rows.shape(0)), + static_cast(local_residual.shape(1)), + static_cast(byte_size(local_residual.dtype(), 8)), + first, + last, + reinterpret_cast(stream_ptr)); + }, + py::arg("local_residual"), + py::arg("selected_local_rows"), + py::arg("candidate_active"), + py::arg("carry"), + py::arg("first"), + py::arg("last"), + py::arg("stream_ptr")); +} + +} // namespace turbomind::python diff --git a/src/turbomind/kernels/gemm/CMakeLists.txt b/src/turbomind/kernels/gemm/CMakeLists.txt index 9e252dca26..4756d0ea01 100644 --- a/src/turbomind/kernels/gemm/CMakeLists.txt +++ b/src/turbomind/kernels/gemm/CMakeLists.txt @@ -157,13 +157,23 @@ add_library(gemm2_core target_link_libraries(gemm2_core PRIVATE parser nvidia::cutlass::cutlass CUDA::cuda_driver) target_compile_definitions(gemm2_core PRIVATE TM_GEMM_HAS_SM90_MIXED=${GEMM2_HAS_SM90_MIXED}) -# Keep all registrations and place their archives before gemm2_core so the -# kernel implementations' core symbols resolve in link order. -list(JOIN GEMM2_KERNEL_TARGETS "," GEMM2_KERNEL_TARGET_LIST) -set_property(TARGET gemm2_core PROPERTY INTERFACE_LINK_LIBRARIES_DIRECT - "$") +# The registrar archives call back into gemm2_core, while gemm2_core needs every +# file-scope Registrar retained. Rescan the whole group to resolve that static +# dependency cycle without relying on consumer link order. +set(_gemm2_link_group "gemm2_core") +foreach(target IN LISTS GEMM2_KERNEL_TARGETS) + string(APPEND _gemm2_link_group ",$") +endforeach() add_library(gemm2 INTERFACE) -target_link_libraries(gemm2 INTERFACE gemm2_core) +target_link_libraries(gemm2 INTERFACE "$") + +# cublasGemmGroupedBatchedEx (CUDA 12.5+): grouped batched GEMM for MoE on SM100 +set(_archs_100 "${CMAKE_CUDA_ARCHITECTURES}") +list(FILTER _archs_100 INCLUDE REGEX "^100") +if(_archs_100 AND CMAKE_CUDA_COMPILER_VERSION VERSION_GREATER_EQUAL "12.5") + target_compile_definitions(gemm2_kernels PRIVATE ENABLE_CUBLAS_GROUPED=1) + message(STATUS "GEMM: ENABLE_CUBLAS_GROUPED=1 (cublasGemmGroupedBatchedEx for MoE on SM100)") +endif() target_compile_options(gemm2_core PRIVATE $<$: diff --git a/src/turbomind/kernels/linear_attn/CMakeLists.txt b/src/turbomind/kernels/linear_attn/CMakeLists.txt index 31de4548b7..3ad01154cf 100644 --- a/src/turbomind/kernels/linear_attn/CMakeLists.txt +++ b/src/turbomind/kernels/linear_attn/CMakeLists.txt @@ -2,6 +2,7 @@ add_subdirectory(kernel) add_library(linear_attn STATIC delta_rule.cu + gdn_state_transaction.cu registry.cu) set_property(TARGET linear_attn PROPERTY POSITION_INDEPENDENT_CODE ON) diff --git a/src/turbomind/kernels/linear_attn/delta_rule.cu b/src/turbomind/kernels/linear_attn/delta_rule.cu index 2a12f97a48..946597a36b 100644 --- a/src/turbomind/kernels/linear_attn/delta_rule.cu +++ b/src/turbomind/kernels/linear_attn/delta_rule.cu @@ -70,6 +70,10 @@ const char* ModeName(GdrMode mode) return "recurrent"; case GdrMode::kChunked: return "chunked"; + case GdrMode::kVerify: + return "verify"; + case GdrMode::kCommit: + return "commit"; } return "invalid"; } @@ -101,7 +105,10 @@ const GdrKernel& RequireGdrKernel(const Plan& plan) bool GatedDeltaRule::Plan(const Operation& requested, const PlanningContext& context, delta_rule::Plan* plan) const { - const bool force_legacy = ForceLegacyGdr(); + const bool force_legacy = ForceLegacyGdr(); + if (force_legacy && (requested.mode == GdrMode::kVerify || requested.mode == GdrMode::kCommit)) { + return false; + } const Operation operation = SelectOperation(requested, force_legacy); const auto architecture = SelectArchitecture(context.arch, force_legacy); const auto& kernel = RequireGdrKernel(operation, context, architecture); @@ -129,4 +136,11 @@ void GatedDeltaRule::Run(const Arguments& args, const delta_rule::Plan& plan, cu RequireGdrKernel(plan).Run(args, plan, stream); } +void GatedDeltaRule::CommitAccepted(const AcceptedPrefixArguments& args, + DataType state_dtype, + cudaStream_t stream) const +{ + invokeCommitAcceptedRecurrentState(args, state_dtype, stream); +} + } // namespace turbomind::linear_attn::delta_rule diff --git a/src/turbomind/kernels/linear_attn/delta_rule.h b/src/turbomind/kernels/linear_attn/delta_rule.h index a2e8f15602..d5ba3b043b 100644 --- a/src/turbomind/kernels/linear_attn/delta_rule.h +++ b/src/turbomind/kernels/linear_attn/delta_rule.h @@ -8,6 +8,7 @@ #include "src/turbomind/core/tensor.h" #include "src/turbomind/kernels/gemm/types.h" +#include "src/turbomind/kernels/linear_attn/gdn_state_transaction.h" namespace turbomind::linear_attn::delta_rule { @@ -16,7 +17,9 @@ class GdrKernel; enum class GdrMode { kRecurrent, - kChunked + kChunked, + kVerify, + kCommit }; constexpr int kAutoGdrChunkSize = 0; @@ -66,12 +69,16 @@ struct PlanningContext { struct Arguments { core::Tensor q, k, v, g, beta; core::Tensor state_ptrs, state_tma_descs, q_offsets, finished; + // Commit consumes [0, commit_lengths[request]) and writes only recurrent state. + // Device int32 [batch], with each length in [0, token_slots]. Q and out are unused. + core::Tensor commit_lengths; core::Tensor* out{}; core::Tensor* workspace{}; int64_t state_layer_offset{}; }; struct Problem { + GdrMode mode{GdrMode::kChunked}; int arch{}; int sm_count{}; DataType input_dtype{kNull}; @@ -95,12 +102,22 @@ struct Problem { inline bool IsRecurrentGdr(const Problem& problem) noexcept { - return problem.chunk_size == kRecurrentGdrChunkSize; + return problem.mode == GdrMode::kRecurrent; } inline bool IsChunkedGdr(const Problem& problem) noexcept { - return problem.chunk_size > kRecurrentGdrChunkSize; + return problem.mode == GdrMode::kChunked; +} + +inline bool IsVerifyGdr(const Problem& problem) noexcept +{ + return problem.mode == GdrMode::kVerify; +} + +inline bool IsCommitGdr(const Problem& problem) noexcept +{ + return problem.mode == GdrMode::kCommit; } struct TensorPlan { @@ -145,6 +162,8 @@ class GatedDeltaRule { const delta_rule::Plan&, cudaStream_t) const; void Run(const Arguments&, const delta_rule::Plan&, cudaStream_t) const; + + void CommitAccepted(const AcceptedPrefixArguments&, DataType state_dtype, cudaStream_t) const; }; } // namespace turbomind::linear_attn::delta_rule diff --git a/src/turbomind/kernels/linear_attn/gdn_state_transaction.cu b/src/turbomind/kernels/linear_attn/gdn_state_transaction.cu new file mode 100644 index 0000000000..9fbe6dfec8 --- /dev/null +++ b/src/turbomind/kernels/linear_attn/gdn_state_transaction.cu @@ -0,0 +1,391 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#include "src/turbomind/kernels/linear_attn/gdn_state_transaction.h" + +#include "src/turbomind/core/data_type.h" +#include "src/turbomind/utils/cuda_utils.h" + +#include +#include + +namespace turbomind::linear_attn::delta_rule { +namespace { + +__global__ void BuildGdnStateStoreMaskKernel( + bool* out, const bool* finished, const bool* speculative, int count) +{ + const int i = blockIdx.x * blockDim.x + threadIdx.x; + if (i < count) { + out[i] = finished[i] || speculative[i]; + } +} + +struct CaptureStrides { + int64_t raw_token; + int64_t key_token; + int64_t key_head; + int64_t value_token; + int64_t value_head; + int64_t log_decay_token; + int64_t beta_token; +}; + +template +__global__ void CaptureGdnTransitionsKernel(const T* __restrict__ raw_projection, + const T* __restrict__ normalized_key, + const T* __restrict__ value, + const float* __restrict__ log_decay, + const float* __restrict__ beta, + const int* __restrict__ q_offsets, + const int* __restrict__ speculative_request_indices, + T* __restrict__ journal_raw, + T* __restrict__ journal_key, + T* __restrict__ journal_value, + float* __restrict__ journal_decay, + float* __restrict__ journal_beta, + CaptureStrides strides, + int speculative_count, + int verify_positions, + int conv_dim, + int hq, + int hv) +{ + const int compact = blockIdx.x; + const int position = blockIdx.y; + const int request = speculative_request_indices[compact]; + const int source_row = q_offsets[request] + position; + const int journal_row = blockIdx.z * speculative_count * verify_positions + + compact * verify_positions + position; + + const T* raw_row = raw_projection + int64_t(source_row) * strides.raw_token; + T* dst_raw = journal_raw + int64_t(journal_row) * conv_dim; + for (int col = threadIdx.x; col < conv_dim; col += blockDim.x) { + dst_raw[col] = raw_row[col]; + } + + const T* key_row = normalized_key + int64_t(source_row) * strides.key_token; + T* dst_key = journal_key + int64_t(journal_row) * hq * 128; + for (int col = threadIdx.x; col < hq * 128; col += blockDim.x) { + const int head = col / 128; + const int dim = col % 128; + dst_key[col] = key_row[int64_t(head) * strides.key_head + dim]; + } + + const T* value_row = value + int64_t(source_row) * strides.value_token; + T* dst_value = journal_value + int64_t(journal_row) * hv * 128; + for (int col = threadIdx.x; col < hv * 128; col += blockDim.x) { + const int head = col / 128; + const int dim = col % 128; + dst_value[col] = value_row[int64_t(head) * strides.value_head + dim]; + } + + const float* decay_row = log_decay + int64_t(source_row) * strides.log_decay_token; + const float* beta_row = beta + int64_t(source_row) * strides.beta_token; + float* dst_decay = journal_decay + int64_t(journal_row) * hv; + float* dst_beta = journal_beta + int64_t(journal_row) * hv; + for (int head = threadIdx.x; head < hv; head += blockDim.x) { + dst_decay[head] = decay_row[head]; + dst_beta[head] = beta_row[head]; + } +} + +template +__global__ void CommitAcceptedConvStateKernel(const T* __restrict__ raw_conv, + void* const* __restrict__ conv_state_ptrs, + const int* __restrict__ request_indices, + const int* __restrict__ entry_sequence_length, + const int* __restrict__ accept_len, + const int* __restrict__ conv_state_offsets, + int speculative_count, + int position_count, + int conv_dim, + int d_conv) +{ + const int channel = blockIdx.x * blockDim.x + threadIdx.x; + const int compact = blockIdx.y; + const int layer = blockIdx.z; + if (channel >= conv_dim) { + return; + } + + const int request = request_indices[compact]; + const int entry = entry_sequence_length[request]; + T* state = static_cast(conv_state_ptrs[request]) + conv_state_offsets[layer]; + const T* journal = raw_conv + + (int64_t(layer) * speculative_count + compact) * position_count * conv_dim; + const int accepted = accept_len[request]; + for (int position = 0; position < accepted; ++position) { + const int ring = (entry - 1 + position) % d_conv; + state[int64_t(ring) * conv_dim + channel] = journal[int64_t(position) * conv_dim + channel]; + } +} + +template +__device__ __forceinline__ float ToFloat(T value) +{ + return static_cast(value); +} + +template +__device__ __forceinline__ T FromFloat(float value) +{ + return static_cast(value); +} + +template +__global__ void CommitAcceptedRecurrentStateKernel(const InputT* __restrict__ key, + const InputT* __restrict__ value, + const float* __restrict__ log_decay, + const float* __restrict__ beta, + void* const* __restrict__ recurrent_state_ptrs, + const int* __restrict__ request_indices, + const int* __restrict__ accept_len, + int64_t layer_group_stride, + int64_t request_stride, + int layer_count, + int speculative_count, + int position_count, + int hq, + int hv, + int num_head_groups, + int layers_per_block, + int heads_per_block, + int total_work) +{ + constexpr int kHeadDim = 128; + constexpr int kTileK = 16; + constexpr int kTileV = 4; + constexpr int kKeyThreads = kHeadDim / kTileK; + + const int key_lane = threadIdx.x % kKeyThreads; + const int value_lane = threadIdx.x / kKeyThreads; + + for (int work = blockIdx.x; work < total_work; work += gridDim.x) { + int index = work; + const int value_head = index % hv; + index /= hv; + const int compact = index % speculative_count; + const int layer = index / speculative_count; + const int request = request_indices[compact]; + + const int layer_group = layer / layers_per_block; + const int layer_in_group = layer % layers_per_block; + const int head_group = value_head / heads_per_block; + const int local_head = value_head % heads_per_block; + void* part = recurrent_state_ptrs[int64_t(layer_group) * layer_group_stride + + int64_t(request) * request_stride + head_group]; + StateT* state = static_cast(part) + + (int64_t(layer_in_group) * heads_per_block + local_head) * kHeadDim * kHeadDim; + + float fragment[kTileK][kTileV]; +#pragma unroll + for (int ki = 0; ki < kTileK; ++ki) { +#pragma unroll + for (int vi = 0; vi < kTileV; ++vi) { + const int key_index = key_lane * kTileK + ki; + const int value_index = value_lane * kTileV + vi; + fragment[ki][vi] = ToFloat(state[int64_t(key_index) * kHeadDim + value_index]); + } + } + + const int accepted = accept_len[request]; + const int key_head = value_head / (hv / hq); + for (int position = 0; position < accepted; ++position) { + const int64_t row = (int64_t(layer) * speculative_count + compact) * position_count + position; + const InputT* key_row = key + (row * hq + key_head) * kHeadDim; + const InputT* value_row = value + (row * hv + value_head) * kHeadDim; + const float decay = exp2f(log_decay[row * hv + value_head] * 1.4426950408889634f); + const float beta_value = beta[row * hv + value_head]; + + float prediction[kTileV]{}; +#pragma unroll + for (int ki = 0; ki < kTileK; ++ki) { + const float key_value = ToFloat(key_row[key_lane * kTileK + ki]); +#pragma unroll + for (int vi = 0; vi < kTileV; ++vi) { + fragment[ki][vi] *= decay; + prediction[vi] += fragment[ki][vi] * key_value; + } + } +#pragma unroll + for (int offset = 4; offset > 0; offset >>= 1) { +#pragma unroll + for (int vi = 0; vi < kTileV; ++vi) { + prediction[vi] += __shfl_xor_sync(0xffffffffu, prediction[vi], offset); + } + } + + float delta[kTileV]; +#pragma unroll + for (int vi = 0; vi < kTileV; ++vi) { + delta[vi] = (ToFloat(value_row[value_lane * kTileV + vi]) - prediction[vi]) * beta_value; + } +#pragma unroll + for (int ki = 0; ki < kTileK; ++ki) { + const float key_value = ToFloat(key_row[key_lane * kTileK + ki]); +#pragma unroll + for (int vi = 0; vi < kTileV; ++vi) { + fragment[ki][vi] += key_value * delta[vi]; + } + } + } + +#pragma unroll + for (int ki = 0; ki < kTileK; ++ki) { +#pragma unroll + for (int vi = 0; vi < kTileV; ++vi) { + const int key_index = key_lane * kTileK + ki; + const int value_index = value_lane * kTileV + vi; + state[int64_t(key_index) * kHeadDim + value_index] = FromFloat(fragment[ki][vi]); + } + } + } +} + +} // namespace + +void invokeBuildGdnStateStoreMask(bool* suppress_state_store, + const bool* finished, + const bool* speculative_row, + int request_count, + cudaStream_t stream) +{ + if (request_count == 0) { + return; + } + const int block = 256; + const int grid = (request_count + block - 1) / block; + BuildGdnStateStoreMaskKernel<<>>( + suppress_state_store, finished, speculative_row, request_count); + TM_CUDA_CHECK(cudaGetLastError()); +} + +void invokeCaptureGdnTransitions(const Tensor& raw_projection, + const Tensor& normalized_key, + const Tensor& value, + const Tensor& log_decay, + const Tensor& beta, + const Buffer_& q_offsets, + const Buffer_& speculative_request_indices, + int gdn_layer, + int verify_positions, + TransitionJournal journal, + cudaStream_t stream) +{ + const int speculative_count = speculative_request_indices.size(); + const int conv_dim = raw_projection.shape(1); + const int hq = normalized_key.shape(2); + const int hv = value.shape(2); + const CaptureStrides strides{raw_projection.stride(0), + normalized_key.stride(1), + normalized_key.stride(2), + value.stride(1), + value.stride(2), + log_decay.stride(1), + beta.stride(1)}; + const dim3 grid(speculative_count, verify_positions, 1); + auto launch = [&](auto type) { + using T = decltype(type); + CaptureGdnTransitionsKernel<<>>(raw_projection.data(), + normalized_key.data(), + value.data(), + log_decay.data(), + beta.data(), + q_offsets.data(), + speculative_request_indices.data(), + journal.raw_conv.data() + + int64_t(gdn_layer) * speculative_count + * verify_positions * conv_dim, + journal.key.data() + + int64_t(gdn_layer) * speculative_count + * verify_positions * hq * 128, + journal.value.data() + + int64_t(gdn_layer) * speculative_count + * verify_positions * hv * 128, + journal.log_decay.data() + + int64_t(gdn_layer) * speculative_count + * verify_positions * hv, + journal.beta.data() + + int64_t(gdn_layer) * speculative_count + * verify_positions * hv, + strides, + speculative_count, + verify_positions, + conv_dim, + hq, + hv); + }; + TM_DISPATCH_DTYPES(raw_projection.dtype(), launch, half_t, bfloat16_t); + TM_CUDA_CHECK(cudaGetLastError()); +} + +void invokeCommitAcceptedConvState(const Tensor& raw_conv, + const Buffer_& conv_state_ptrs, + const Buffer_& speculative_request_indices, + const Buffer_& entry_sequence_length, + const Buffer_& accept_len, + const Buffer_& conv_state_offsets, + int conv_dim, + int d_conv, + cudaStream_t stream) +{ + const int layers = raw_conv.shape(0); + const int speculative_count = raw_conv.shape(1); + const int position_count = raw_conv.shape(2); + const dim3 grid((conv_dim + 255) / 256, speculative_count, layers); + auto launch = [&](auto type) { + using T = decltype(type); + CommitAcceptedConvStateKernel<<>>(raw_conv.data(), + conv_state_ptrs.data(), + speculative_request_indices.data(), + entry_sequence_length.data(), + accept_len.data(), + conv_state_offsets.data(), + speculative_count, + position_count, + conv_dim, + d_conv); + }; + TM_DISPATCH_DTYPES(raw_conv.dtype(), launch, half_t, bfloat16_t); + TM_CUDA_CHECK(cudaGetLastError()); +} + +void invokeCommitAcceptedRecurrentState( + const AcceptedPrefixArguments& args, DataType state_dtype, cudaStream_t stream) +{ + const int total_work = args.layer_count * args.speculative_count * args.hv; + const int grid = std::min(total_work, args.sm_count * 4); + const int64_t layer_group_stride = args.recurrent_state_ptrs.stride(0); + const int64_t request_stride = args.recurrent_state_ptrs.stride(1); + + auto launch_input = [&](auto input_type) { + using InputT = decltype(input_type); + auto launch_state = [&](auto state_type) { + using StateT = decltype(state_type); + CommitAcceptedRecurrentStateKernel<<>>( + args.key.data(), + args.value.data(), + args.log_decay.data(), + args.beta.data(), + args.recurrent_state_ptrs.data(), + args.request_indices.data(), + args.accept_len.data(), + layer_group_stride, + request_stride, + args.layer_count, + args.speculative_count, + args.position_count, + args.hq, + args.hv, + args.num_head_groups, + args.layers_per_block, + args.heads_per_block, + total_work); + }; + TM_DISPATCH_DTYPES(state_dtype, launch_state, bfloat16_t, float); + }; + TM_DISPATCH_DTYPES(args.key.dtype(), launch_input, half_t, bfloat16_t); + TM_CUDA_CHECK(cudaGetLastError()); +} + +} // namespace turbomind::linear_attn::delta_rule diff --git a/src/turbomind/kernels/linear_attn/gdn_state_transaction.h b/src/turbomind/kernels/linear_attn/gdn_state_transaction.h new file mode 100644 index 0000000000..efc58236cb --- /dev/null +++ b/src/turbomind/kernels/linear_attn/gdn_state_transaction.h @@ -0,0 +1,73 @@ +// Copyright (c) OpenMMLab. All rights reserved. +#pragma once + +#include "src/turbomind/core/buffer.h" +#include "src/turbomind/core/tensor.h" +#include "src/turbomind/kernels/gemm/types.h" + +#include + +namespace turbomind::linear_attn::delta_rule { + +using core::Buffer_; +using core::Tensor; + +struct TransitionJournal { + Tensor raw_conv; + Tensor key; + Tensor value; + Tensor log_decay; + Tensor beta; +}; + +void invokeBuildGdnStateStoreMask(bool* suppress_state_store, + const bool* finished, + const bool* speculative_row, + int request_count, + cudaStream_t stream); + +void invokeCaptureGdnTransitions(const Tensor& raw_projection, + const Tensor& normalized_key, + const Tensor& value, + const Tensor& log_decay, + const Tensor& beta, + const Buffer_& q_offsets, + const Buffer_& speculative_request_indices, + int gdn_layer, + int verify_positions, + TransitionJournal journal, + cudaStream_t stream); + +void invokeCommitAcceptedConvState(const Tensor& raw_conv, + const Buffer_& conv_state_ptrs, + const Buffer_& speculative_request_indices, + const Buffer_& entry_sequence_length, + const Buffer_& accept_len, + const Buffer_& conv_state_offsets, + int conv_dim, + int d_conv, + cudaStream_t stream); + +struct AcceptedPrefixArguments { + Tensor key; + Tensor value; + Tensor log_decay; + Tensor beta; + Tensor recurrent_state_ptrs; + Tensor request_indices; + Tensor accept_len; + int layer_count{}; + int speculative_count{}; + int position_count{}; + int hq{}; + int hv{}; + int num_head_groups{}; + int layers_per_block{}; + int heads_per_block{}; + int sm_count{}; +}; + +void invokeCommitAcceptedRecurrentState( + const AcceptedPrefixArguments& args, DataType state_dtype, cudaStream_t stream); + +} // namespace turbomind::linear_attn::delta_rule diff --git a/src/turbomind/kernels/linear_attn/kernel/CMakeLists.txt b/src/turbomind/kernels/linear_attn/kernel/CMakeLists.txt index 0da5247a33..44836579a9 100644 --- a/src/turbomind/kernels/linear_attn/kernel/CMakeLists.txt +++ b/src/turbomind/kernels/linear_attn/kernel/CMakeLists.txt @@ -9,7 +9,8 @@ set(LINEAR_ATTN_GDR_SM90_KERNEL_SOURCES sm_90/plan.cc sm_90/tma_desc_prepare.cu sm_90/kkt_solve.cu - sm_90/recurrent.cu) + sm_90/recurrent.cu + sm_90/verification_fwd.cu) set(LINEAR_ATTN_GDR_SM120_KERNEL_SOURCES sm_120/entry.cu diff --git a/src/turbomind/kernels/linear_attn/kernel/plan.cc b/src/turbomind/kernels/linear_attn/kernel/plan.cc index 5c6f9912e2..d6c3ccd7f3 100644 --- a/src/turbomind/kernels/linear_attn/kernel/plan.cc +++ b/src/turbomind/kernels/linear_attn/kernel/plan.cc @@ -41,6 +41,7 @@ Problem BuildProblem(const PlanningContext& context, const GdrKernelSpec& spec) { Problem problem{}; problem.arch = context.arch; + problem.mode = spec.mode; problem.sm_count = context.sm_count; problem.input_dtype = context.input_dtype; problem.state_dtype = context.state_dtype; @@ -56,7 +57,7 @@ Problem BuildProblem(const PlanningContext& context, const GdrKernelSpec& spec) problem.chunk_size = spec.chunk_size; problem.num_head_groups = context.num_head_groups; problem.heads_per_block = context.heads_per_block; - if (spec.mode == GdrMode::kRecurrent) { + if (spec.mode == GdrMode::kRecurrent || spec.mode == GdrMode::kVerify || spec.mode == GdrMode::kCommit) { problem.sequence_num = context.physical_batch; problem.total_chunks = context.physical_batch; problem.max_sequence_chunks = context.physical_batch > 0 ? 1 : 0; @@ -105,6 +106,16 @@ void BuildOptimizedTensorPlans(Plan* plan, size_t direct_descriptor_bytes) TensorPlan{core::Layout{{problem.batch, problem.token_num, problem.hv, 128}, {value_batch, value_row, 128, 1}}, problem.input_dtype, value_elements}; + if (IsCommitGdr(problem)) { + plan->out = TensorPlan{core::Layout{{0}}, problem.input_dtype}; + } + if (IsVerifyGdr(problem) || IsCommitGdr(problem)) { + plan->g_cumsum = TensorPlan{core::Layout{{0}}, kFloat32}; + plan->resolvent = TensorPlan{core::Layout{{0}}, problem.input_dtype}; + plan->workspace = TensorPlan{core::Layout{{0}}, kUint8}; + plan->workspace_bytes = 0; + return; + } const bool chunked = IsChunkedGdr(problem); const core::ssize_t gate_stride = chunked ? core::ssize_t(AlignUp(size_t(problem.hv), 4)) : problem.gate_stride; const core::ssize_t gate_batch_stride = diff --git a/src/turbomind/kernels/linear_attn/kernel/sm_90/internal.h b/src/turbomind/kernels/linear_attn/kernel/sm_90/internal.h index 057f6a7782..6f7d2315f3 100644 --- a/src/turbomind/kernels/linear_attn/kernel/sm_90/internal.h +++ b/src/turbomind/kernels/linear_attn/kernel/sm_90/internal.h @@ -116,6 +116,8 @@ void LaunchSm90Recurrent(const core::Tensor&, DataType, cudaStream_t); void PrepareSm90RecurrentStateTmaDescriptors(const core::Tensor&, core::Tensor&, int, int, const Plan&, cudaStream_t); +void PrepareSm90StateTmaDescriptors( + const core::Tensor&, core::Tensor&, int, int, int, const CUtensorMap&, cudaStream_t); void LaunchSm90KktSolve(const core::Tensor&, const core::Tensor&, const core::Tensor&, diff --git a/src/turbomind/kernels/linear_attn/kernel/sm_90/plan.cc b/src/turbomind/kernels/linear_attn/kernel/sm_90/plan.cc index ee3ede4db1..b92afbff34 100644 --- a/src/turbomind/kernels/linear_attn/kernel/sm_90/plan.cc +++ b/src/turbomind/kernels/linear_attn/kernel/sm_90/plan.cc @@ -163,7 +163,7 @@ bool PlanSm90Operation(const GdrKernelSpec& spec, KktTmaDescriptorBytes(plan->problem) + FusedGdrTmaDescriptorBytes(plan->problem) + 127; BuildOptimizedTensorPlans(plan, descriptor_bytes); plan->state_tma_desc_bytes_per_layer_group = - IsRecurrentGdr(plan->problem) ? + IsRecurrentGdr(plan->problem) || IsVerifyGdr(plan->problem) || IsCommitGdr(plan->problem) ? size_t(plan->problem.sequence_num) * plan->problem.num_head_groups * sizeof(CUtensorMap) : 0; return true; diff --git a/src/turbomind/kernels/linear_attn/kernel/sm_90/recurrent.cu b/src/turbomind/kernels/linear_attn/kernel/sm_90/recurrent.cu index 0e648fff11..3d3f946172 100644 --- a/src/turbomind/kernels/linear_attn/kernel/sm_90/recurrent.cu +++ b/src/turbomind/kernels/linear_attn/kernel/sm_90/recurrent.cu @@ -1087,7 +1087,6 @@ int SelectRecurrentGdrBlockDv(const Problem& problem, DataType state_dtype) return SelectRecurrentGdrBlockDv<__nv_bfloat16>(problem); } -template __global__ __launch_bounds__(32, 1) void PrepareGroupedStateDescriptors(const __grid_constant__ CUtensorMap state_tma_desc, const int64_t* addresses, @@ -1112,7 +1111,7 @@ __global__ __launch_bounds__(32, CopyTmaDescriptor(&smem_descriptor, &state_tma_desc, lane, 32); __syncwarp(); if (lane == 0) { - auto* state_base = reinterpret_cast(static_cast(addresses[pointer_index])); + const void* state_base = reinterpret_cast(static_cast(addresses[pointer_index])); ReplaceTmaAddress(&smem_descriptor, state_base); } __syncwarp(); @@ -1136,18 +1135,8 @@ void PrepareSm90RecurrentStateTmaDescriptorsTyped(const core::Tensor& state_ptrs using Kernel = Sm90GdrRecurrent; const auto state_tma_desc = Kernel::MakeStateTmaDesc( reinterpret_cast(state_tma_descs.raw_data()), layers_per_block, heads_per_block, BlockDv); - const auto* addresses = reinterpret_cast(state_ptrs.raw_data()); - auto* descriptors = reinterpret_cast(state_tma_descs.raw_data()); - const int work = layer_groups * sequence_count * num_head_groups; - PrepareGroupedStateDescriptors<<>>(state_tma_desc, - addresses, - state_ptrs.stride(0), - state_ptrs.stride(1), - state_ptrs.stride(2), - descriptors, - sequence_count, - num_head_groups); - TM_CUDA_CHECK(cudaGetLastError()); + detail::PrepareSm90StateTmaDescriptors( + state_ptrs, state_tma_descs, layer_groups, sequence_count, num_head_groups, state_tma_desc, stream); } template @@ -1218,6 +1207,28 @@ void LaunchSm90GdrRecurrentTyped(const core::Tensor& q, namespace detail { +void PrepareSm90StateTmaDescriptors(const core::Tensor& state_ptrs, + core::Tensor& state_tma_descs, + int layer_groups, + int sequence_count, + int num_head_groups, + const CUtensorMap& prototype, + cudaStream_t stream) +{ + const auto* addresses = reinterpret_cast(state_ptrs.raw_data()); + auto* descriptors = reinterpret_cast(state_tma_descs.raw_data()); + const int work = layer_groups * sequence_count * num_head_groups; + PrepareGroupedStateDescriptors<<>>(prototype, + addresses, + state_ptrs.stride(0), + state_ptrs.stride(1), + state_ptrs.stride(2), + descriptors, + sequence_count, + num_head_groups); + TM_CUDA_CHECK(cudaGetLastError()); +} + void LaunchSm90Recurrent(const core::Tensor& q, const core::Tensor& k, const core::Tensor& v, diff --git a/src/turbomind/kernels/linear_attn/kernel/sm_90/verification_fwd.cu b/src/turbomind/kernels/linear_attn/kernel/sm_90/verification_fwd.cu new file mode 100644 index 0000000000..f4f91c4d47 --- /dev/null +++ b/src/turbomind/kernels/linear_attn/kernel/sm_90/verification_fwd.cu @@ -0,0 +1,2261 @@ +#include "src/turbomind/kernels/linear_attn/kernel/sm_90/internal.h" + +#include "src/turbomind/kernels/linear_attn/kernel/plan.h" +#include "src/turbomind/kernels/linear_attn/kernel/sm_90/common.h" +#include "src/turbomind/kernels/linear_attn/registrar.h" +#include "src/turbomind/utils/cuda_utils.h" + +#include +#include + +#include +#include +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace turbomind::linear_attn::delta_rule { +namespace { + +using namespace cute; + +template +__global__ __launch_bounds__(Operator::MaxThreadsPerBlock, Operator::MinBlocksPerMultiprocessor) +void Sm90GdrVerifyCommitDeviceKernel(const __grid_constant__ typename Operator::Params parameters) +{ + extern __shared__ __align__(Operator::SharedMemoryAlignment) unsigned char shared_bytes[]; + auto& shared_storage = *reinterpret_cast(shared_bytes); + Operator{}(parameters, shared_storage); +} + +// verify: +// +// WG-0 WG-1 WG-2 +// TC gram=_gram_(K) TC P=_p_(Q,K) prefix=_prefix_(gate) +// AG=_ag_(gram) TC O0=_o_sq_(S,Q) TC U=_u_(S,K) +// _ag_scale_(prefix) W=_w_(V,U,prefix) +// TC AGW=_agw_(W,AG) +// O1=_o_scale_(O0,prefix) +// _p_scale_(P,prefix) +// TC O2=_o_final_(O1,AGW,P) +// _store_(O2) +// +// commit: +// +// WG-0 WG-1 WG-2 +// TC gram=_gram_(K) prefix=_prefix_(gate) +// AG=_ag_(gram) TC U=_u_(S,K) +// _ag_scale_(prefix) W=_w_(V,U,prefix) +// TC AGW=_agw_(W,AG) +// Kc=_k_scale_(K,prefix,L) +// TC S'=_commit_(S,AGW,Kc) +// _store_(S') +// + +template +class Sm90GdrVerifyCommitKernel final: public GdrKernel { +private: + static_assert(Mode == GdrMode::kVerify || Mode == GdrMode::kCommit); + static constexpr bool kCommit = Mode == GdrMode::kCommit; + // Unqualified bfloat16_t and gemm resolve against TurboMind declarations; + // qualify only these compiler-ambiguous CuTe symbols. + using Element = cute::bfloat16_t; + using MmaAtom = MMA_Atom; + + static constexpr int kHeadDim = 128; + static constexpr int kBlockDv = 128; + static constexpr int kHalfDv = 64; + static constexpr int kMmaM = 16; + static constexpr int kMmaN = 8; + static constexpr int kWarpThreads = 32; + static constexpr int kComputeWarpGroups = 3; + static constexpr int kWarpGroups = kComputeWarpGroups + 1; + static constexpr int kTmaGlobalAddressAlignment = 16; + static constexpr int kTmaWgmmaSmemAlignment = 1024; + static constexpr int kBarrierAlignment = 16; + static constexpr int kValueHalves = kBlockDv / kHalfDv; + static constexpr int kCapacityNTiles = Capacity / kMmaN; + static constexpr float kHeadScale = 0.08838834764831845f; + static constexpr int kStateElementsPerHead = kHeadDim * kBlockDv; + static constexpr int kWStageBytes = kBlockDv * kMmaM * sizeof(Element); + static constexpr int kRemovedQKVTailSharedBytesPerStage = + (kMmaM - Capacity) * (2 * kHeadDim + kBlockDv) * sizeof(Element); + static constexpr int kRemovedGateExpDifferenceTailSharedBytesPerStage = + (kMmaM * kMmaM - Capacity * Capacity) * sizeof(float); + + struct Fp32StateTraits { + using State = float; + using SmemSwizzle = Swizzle<3, 4, 3>; + + static constexpr DataType kDataType = kFloat32; + static constexpr int kStages = 2; + static constexpr int kStateOperandOffsetBytes = kStages * kStateElementsPerHead * sizeof(State); + static constexpr int kStateOperandStageBytes = 0; + static constexpr int kStateStorageBytes = + kStateOperandOffsetBytes + kStateElementsPerHead * sizeof(Element); + static constexpr int kPersistentSharedBytes = + 203776 + 3 * kTmaWgmmaSmemAlignment + 2 * kWStageBytes + - kRemovedQKVTailSharedBytesPerStage * kStages + - kRemovedGateExpDifferenceTailSharedBytesPerStage * (kStages + 1) + / kTmaWgmmaSmemAlignment * kTmaWgmmaSmemAlignment; + static constexpr const char* kName8 = + kCommit ? "sm90_delta_rule_commit8_f32_state" : "sm90_delta_rule_verify8_f32_state"; + static constexpr const char* kName16 = + kCommit ? "sm90_delta_rule_commit16_f32_state" : "sm90_delta_rule_verify16_f32_state"; + }; + + struct Bf16StateTraits { + using State = Element; + using SmemSwizzle = Swizzle<2, 4, 3>; + + static constexpr DataType kDataType = kBfloat16; + static constexpr int kStages = 3; + static constexpr int kStateOperandOffsetBytes = 0; + static constexpr int kStateOperandStageBytes = kStateElementsPerHead * sizeof(State); + static constexpr int kStateStorageBytes = kStages * kStateOperandStageBytes; + static constexpr int kPersistentSharedBytes = + 159744 + kTmaWgmmaSmemAlignment + 2 * kWStageBytes + - kRemovedQKVTailSharedBytesPerStage * kStages + - kRemovedGateExpDifferenceTailSharedBytesPerStage * (kStages + 1) + / kTmaWgmmaSmemAlignment * kTmaWgmmaSmemAlignment; + static constexpr const char* kName8 = + kCommit ? "sm90_delta_rule_commit8_bf16_state" : "sm90_delta_rule_verify8_bf16_state"; + static constexpr const char* kName16 = + kCommit ? "sm90_delta_rule_commit16_bf16_state" : "sm90_delta_rule_verify16_bf16_state"; + }; + + static_assert(StateType == kFloat32 || StateType == kBfloat16); + using StateTraits = std::conditional_t; + using StateT = typename StateTraits::State; + using State = StateT; + static constexpr int kStages = StateTraits::kStages; + static constexpr int kHandoffStages = kStages + 1; + static constexpr int kInvalidRequest = -1; + static constexpr unsigned kWarpMask = 0xffffffffu; + static constexpr int kMinimumTokens = kCommit ? 1 : 2; + static constexpr int kLowerTriangularBankPadding = kMmaN / 2; + static constexpr int kLowerTriangularRowStride = kMmaM + kLowerTriangularBankPadding; + static constexpr int kHardwareNamedBarrierCount = 16; + + enum class NamedBarrierId : int { + PReady = 0, + GramReady, + WReady, + AGReady = WReady + kHandoffStages, + StateDv0Ready = AGReady + kHandoffStages, + StateDv1Ready, + GatePrefixReady, + AGSolveReady, + }; + + enum class WarpGroupRole : int { + GramPAG, + Output, + Update, + Producer, + }; + + static_assert(Capacity == kMmaN || Capacity == kMmaM); + static_assert(kWarpThreads % kMmaM == 0); + static_assert(static_cast(NamedBarrierId::AGSolveReady) + < kHardwareNamedBarrierCount); + + using HeadDim = Int; + using BlockDv = Int; + using HalfDv = Int; + using MmaM = Int; + using MmaN = Int; + using WarpThreads = Int; + using AGBlock = _4; + using AGBlocks = Int; + using QKVector = HalfDv; + using QKVectors = Int; + using StateDvBoxes = Int; + using StateDkPanels = Int; + using GateLaneRows = Int; + using GateLaneLayout = Layout>; + using ValueDvCoordinateLayout = Layout>>; + + using SmemLayoutQKTile = decltype( + tile_to_shape(GMMA::Layout_K_SW128_Atom{}, + Shape, HeadDim>{}, + Step<_1, _2>{})); + using SmemLayoutQKMmaOperandTile = decltype( + tile_to_shape(GMMA::Layout_K_SW128_Atom{}, + Shape{}, + Step<_1, _2>{})); + using SmemBf16PointerFlagBits = smem_ptr_flag_bits::value>; + using SmemQKSwizzle = decltype(get_swizzle_portion(SmemLayoutQKTile{})); + using SmemLayoutQKTileLinear = decltype(get_nonswizzle_portion(SmemLayoutQKTile{})); + using SmemLayoutQKTmaTileLinear = decltype(select<1, 3, 2>( + flatten(zipped_divide(SmemLayoutQKTileLinear{}, Tile<_1, QKVector>{})))); + using SmemLayoutQKTmaTile = decltype(composition( + SmemQKSwizzle{}, SmemBf16PointerFlagBits{}, SmemLayoutQKTmaTileLinear{})); + using SmemLayoutValueTma = decltype( + tile_to_shape( + GMMA::Layout_MN_SW128_Atom{}, Shape>{})); + using Stages = Int; + using HandoffStages = Int; + using SmemLayoutAG = decltype(tile_to_shape( + GMMA::Layout_K_SW32_Atom{}, + Shape{}, + Step<_1, _2, _3>{})); + using SmemLayoutP = decltype(tile_to_shape( + GMMA::Layout_K_SW32_Atom{}, + Shape{}, + Step<_1, _2>{})); + using SmemLayoutW = decltype(tile_to_shape( + GMMA::Layout_MN_SW128_Atom{}, + Shape{}, + Step<_1, _2, _3>{})); + using SmemLayoutAGW = decltype(tile_to_shape( + GMMA::Layout_MN_SW128_Atom{}, + Shape{}, + Step<_1, _2>{})); + using WgmmaBOperandTailCoordinates = + Layout, MmaN>, Stride>; + using LowerTriangularRowStride = Int; + using SmemLayoutLowerTriangularTile = + Layout, Stride>; + static constexpr int kLowerTriangularElements = cosize_v; + using SmemLayoutLowerTriangular = + Layout, + Stride>>; + using SmemLayoutBetaCheckpoint = + Layout, Stride<_1, MmaM>>; + using SmemLayoutGateExp = + Layout, Stride<_1, MmaM>>; + using SmemLayoutGateExpDifferenceStage = decltype(tile_to_shape( + Layout, Stride>{}, + Shape, Int>{}, + Step<_1, _2>{})); + using SmemLayoutGateExpDifference = decltype(tile_to_shape( + SmemLayoutGateExpDifferenceStage{}, + Shape, Int, HandoffStages>{}, + Step<_1, _2, _3>{})); + static_assert(kLowerTriangularElements >= kMmaM * kMmaM); + + static constexpr int kOutputN = Capacity; + static constexpr int kThreads = kWarpGroups * cutlass::NumThreadsPerWarpGroup; + static constexpr int kComputeThreads = kComputeWarpGroups * cutlass::NumThreadsPerWarpGroup; + static constexpr int kSVConsumerThreads = + kComputeThreads - cutlass::NumThreadsPerWarpGroup; + static constexpr int kComputeWarps = kComputeThreads / kWarpThreads; + static constexpr int kProducerWarp = kComputeWarps; + static constexpr int kProducerThread = kComputeThreads; + static constexpr int kProducerRegisters = 48; + static constexpr int kGramPAGRegisters = 144; + static constexpr int kOutputRegisters = 160; + static constexpr int kUpdateRegisters = 160; + static constexpr int kRegisterBudgetPerThread = 128; + static constexpr int kGramWarpBegin = 0; + static constexpr int kAGWarp = kGramWarpBegin; + static constexpr int kGramWarps = kMmaM / kMmaN; + static constexpr int kGramPAGWarps = cutlass::NumWarpsPerWarpGroup; + static constexpr int kOutputHeadCheckpointWarp = cutlass::NumWarpsPerWarpGroup + 2; + static constexpr int kPrefixWarp = kComputeWarps - 1; + static constexpr int kPrefixRowsPerWarp = Capacity / cutlass::NumWarpsPerWarpGroup; + static constexpr int kPrefixColumnVectors = Capacity / _4::value; + static constexpr int kQKBytes = Capacity * kHeadDim * sizeof(Element); + static constexpr int kPayloadValueBytes = Capacity * kBlockDv * sizeof(Element); + static constexpr int kStateBytesPerHead = kStateElementsPerHead * sizeof(State); + static constexpr int kQKTransactionBytes = (kCommit ? 1 : 2) * kQKBytes; + static constexpr int kSVTransactionBytes = kPayloadValueBytes + kStateBytesPerHead; + static_assert(Capacity % AGBlock::value == 0); + static_assert(AGBlocks::value <= cutlass::NumWarpsPerWarpGroup); + static_assert(kGramWarpBegin + kCapacityNTiles <= kGramWarps); + static_assert(kAGWarp < kGramPAGWarps); + static_assert(kCapacityNTiles <= cutlass::NumWarpsPerWarpGroup); + static_assert(Capacity % cutlass::NumWarpsPerWarpGroup == 0); + static_assert(Capacity % _4::value == 0); + static_assert(kPrefixWarp / cutlass::NumWarpsPerWarpGroup + == static_cast(WarpGroupRole::Update)); + static_assert(kProducerWarp / cutlass::NumWarpsPerWarpGroup + == static_cast(WarpGroupRole::Producer)); + static_assert(kProducerRegisters + kGramPAGRegisters + kOutputRegisters + kUpdateRegisters + <= kWarpGroups * kRegisterBudgetPerThread); + static_assert(kMmaM % kMmaN == 0); + static_assert(kMmaN * kMmaN % kWarpThreads == 0); + static_assert(kMmaM * kMmaM % kWarpThreads == 0); + + using PrefixWarpGroupThreadLayout = + Layout, WarpThreads>, + Stride>; + using PrefixOutputLaneLayout = + Layout, Int>, + Stride<_1, Int>>; + using PrefixOutputRowLayout = + Layout, Int>, + Stride, _1>>; + using PrefixColumnVectorLayout = + Layout, _4>, Stride<_4, _1>>; + static_assert(size(PrefixOutputLaneLayout{}) <= kWarpThreads); + using AGBlockTile = Tile; + using AGBlockTranspose = Layout, Stride>; + using AGBlockMma = TiledMMA>, + Layout>>; + using AGBlockThreads = Int; + static_assert(AGBlockThreads::value <= kWarpThreads); + using AGScaleVectorElements = + std::conditional_t; + using AGScaleColumnVectors = Int; + using AGScaleRowsPerWarp = Int; + using AGScaleLaneCoordinates = + Layout, + Stride<_1, AGScaleRowsPerWarp>>; + using AGScaleRowCoordinates = + Layout, AGScaleRowsPerWarp>, + Stride>; + using AGScaleTile = Tile<_1, AGScaleVectorElements>; + using AGScaleLoadAtom = Copy_Atom< + UniversalCopy::value + * AGScaleVectorElements::value>>, + float>; + using AGScaleStoreAtom = Copy_Atom< + UniversalCopy::value + * AGScaleVectorElements::value>>, + Element>; + static_assert(size(AGScaleLaneCoordinates{}) <= kWarpThreads); + static_assert(size(AGScaleRowCoordinates{}) == Capacity); + using OutputN = Int; + using StateTmaTile = Tile, _1>; + using GmemLayoutState = Layout, + Stride<_1, BlockDv, Int>>; + using GmemLayoutStateTma = decltype(select<0, 1, 3, 4, 5>( + flatten(zipped_divide(GmemLayoutState{}, Tile{})))); + using SmemLayoutStateTmaTileLinear = + Layout>>; + using SmemLayoutStateTmaTile = decltype(composition( + typename StateTraits::SmemSwizzle{}, + smem_ptr_flag_bits::value>{}, + SmemLayoutStateTmaTileLinear{})); + using SmemLayoutStatePanel = decltype(take<0, 3>(SmemLayoutStateTmaTile{})); + using SmemLayoutStatePanels = decltype(append<4>( + SmemLayoutStatePanel{}, + Layout>>{})); + using SmemLayoutState = decltype(append<5>( + SmemLayoutStatePanels{}, + Layout, Int>{})); + using StateMmaLayout = + Layout, Shape>, + Stride>, + Stride>>>; + using SmemLayoutStateOperandStorage = decltype(composition( + Swizzle<2, 4, 3>{}, + SmemBf16PointerFlagBits{}, + Layout>{})); + using SmemLayoutStateOperand = decltype( + SmemLayoutStateOperandStorage{}.compose(StateMmaLayout{})); + using StateConversionThreads = Int; + using StateConversionVectorElements = _4; + using StateConversionTile = Tile; + using StateConversionIterations = + Int; + using StateConversionThreadValueLayout = + Layout>, + Stride>>>; + static_assert(cosize_v == size(StateTmaTile{})); + static_assert(size(StateConversionThreadValueLayout{}) == size(StateConversionTile{})); + + using Policy = Sm90GdrVerifyCommitKernel; + + static constexpr const char* kName = Capacity == kMmaN ? StateTraits::kName8 : StateTraits::kName16; + static constexpr GdrKernelSpec kSpec{ + "sm90", Mode, kBfloat16, StateTraits::kDataType, kHeadDim, Capacity}; + + using QKTmaTile = Tile, _1, _1>; + using ValueTmaTile = Tile>; + using PipelineShape = Shape<_1, _1, _1>; + using QKMmaTileNK = Tile; + static constexpr SM90_TMA_LOAD kTmaLoad{}; + using TmaState = decltype(make_tma_copy( + kTmaLoad, + make_tensor(make_gmem_ptr(static_cast(nullptr)), GmemLayoutStateTma{}), + SmemLayoutStateTmaTile{}, + StateTmaTile{}, + _1{})); + using MmaAtomThrLayout = Layout>; + using WgmmaAtomLayout = Layout>; + using QKTiledMma = decltype(make_tiled_mma(MmaAtom{}, MmaAtomThrLayout{}, Tile{})); + using TokenRows = std::conditional_t< + Capacity == kMmaN, + Coord, + Coord>; + using PAccumulatorStoreOperation = std::conditional_t< + Capacity == kMmaN, + SM90_U32x1_STSM_N, + SM90_U32x2_STSM_N>; + using SSWgmmaTiledMma = decltype(make_tiled_mma( + GMMA::ss_op_selector, + GMMA::Major::MN, + GMMA::Major::K>(), + WgmmaAtomLayout{})); + using QKMmaAValidTVCoordinate = + Coord, Coord>>; + using QKMmaAValidTVLayout = decltype( + QKTiledMma{}.get_layoutA_TV()(QKMmaAValidTVCoordinate{})); + using QKCapacityMmaATVLayout = decltype( + group<1, QKMmaAValidTVLayout::rank>(QKMmaAValidTVLayout{})); + using QKCapacityTiledCopyA = decltype(make_tiled_copy_impl( + Copy_Atom{}, + QKCapacityMmaATVLayout{}, + make_shape(MmaM{}, MmaM{}))); + using QKMmaFullTiledCopyA = decltype(make_tiled_copy_A( + Copy_Atom{}, QKTiledMma{})); + using QKTiledCopyA = std::conditional_t; + using QKCopyOperationA = std::conditional_t; + using QKCopyAtomA = Copy_Atom; + using QKMmaAFragmentMCoordinate = std::conditional_t, + Underscore>; + using SmemCopyAtomBKMajor = Copy_Atom; + using WgmmaAccumulatorStoreOperation = + std::conditional_t; + using WgmmaAccumulatorStoreAtom = Copy_Atom; + using QKTiledCopyBKMajor = decltype(make_tiled_copy_B(SmemCopyAtomBKMajor{}, QKTiledMma{})); + using OutputStoreAtom = Copy_Atom, Element>; + using OutputStoreVectorElements = Int; + using OutputStoreTiledCopy = + decltype(make_tiled_copy_C(OutputStoreAtom{}, SSWgmmaTiledMma{})); + static_assert(OutputStoreAtom::NumValSrc == OutputStoreAtom::NumValDst); + static_assert(size(SSWgmmaTiledMma{}) == cutlass::NumThreadsPerWarpGroup); + + using StateUpdateTiledMma = decltype(make_tiled_mma( + GMMA::ss_op_selector, + GMMA::Major::MN, GMMA::Major::MN>(), + WgmmaAtomLayout{})); + using StateUpdateTile = Tile; + using StateUpdateTiles = Layout, Int>>; + using SmemLayoutCommitKey = decltype(tile_to_shape( + GMMA::Layout_MN_SW128_Atom{}, Shape{})); + using CommitKeyCopyAtom = Copy_Atom, Element>; + using CommitKeyTiledCopy = decltype(make_tiled_copy( + CommitKeyCopyAtom{}, + Layout, Stride<_1, _16>>{}, + Layout>{})); + + using QKPipeline = cutlass::PipelineTmaAsync; + using SVPipeline = cutlass::PipelineTmaAsync; + using QKPipelineState = typename QKPipeline::PipelineState; + using SVPipelineState = typename SVPipeline::PipelineState; + using QKPipelineStorage = typename QKPipeline::SharedStorage; + using SVPipelineStorage = typename SVPipeline::SharedStorage; + using SmemLayoutQK = decltype(tile_to_shape( + GMMA::Layout_K_SW128_Atom{}, + Shape, HeadDim, Int>{}, + Step<_1, _2, _3>{})); + using SmemLayoutQKLinear = decltype(get_nonswizzle_portion(SmemLayoutQK{})); + using SmemLayoutQKTmaLinear = decltype(select<1, 4, 2, 3, 0, 5>( + flatten(zipped_divide(SmemLayoutQKLinear{}, Tile<_1, QKVector, _1>{})))); + using SmemLayoutQKTma = decltype(composition( + SmemQKSwizzle{}, SmemBf16PointerFlagBits{}, SmemLayoutQKTmaLinear{})); + using SmemLayoutValue = decltype(tile_to_shape( + GMMA::Layout_MN_SW128_Atom{}, + Shape, Int, Int>{})); + using SmemLayoutGate = Layout>, + Stride<_1, MmaM, WarpThreads>>; + using SmemLayoutBeta = SmemLayoutGate; + + static_assert(SmemLayoutQK{}(_0{}, _0{}, _1{}) == Capacity * kHeadDim); + static_assert(SmemLayoutValue{}(_0{}, _0{}, _0{}, _1{}) == Capacity * kBlockDv); + static_assert(SmemLayoutState{}(_0{}, _0{}, _0{}, _0{}, _1{}) + == kStateElementsPerHead); + + struct OutputHeadCoordinate { + int request; + int value_head; + int commit_length; + }; + + struct alignas(kTmaWgmmaSmemAlignment) SharedStorage { + using T = typename Policy::Element; + array_aligned state; + array_aligned, + Policy::kTmaWgmmaSmemAlignment> q; + array_aligned, + Policy::kTmaWgmmaSmemAlignment> k; + array_aligned, + Policy::kTmaWgmmaSmemAlignment> value; + array_aligned, + Policy::kBarrierAlignment> gate; + array_aligned, + Policy::kBarrierAlignment> beta; + array_aligned, + Policy::kBarrierAlignment> beta_checkpoint; + array_aligned, + Policy::kBarrierAlignment> gate_exp; + array_aligned, + Policy::kBarrierAlignment> gate_exp_difference; + alignas(Policy::kBarrierAlignment) QKPipelineStorage qk_pipeline; + alignas(Policy::kBarrierAlignment) SVPipelineStorage sv_pipeline; + alignas(Policy::kBarrierAlignment) cutlass::arch::ClusterBarrier output_head_ready[kStages]; + alignas(Policy::kBarrierAlignment) uint64_t prefix_ready[kHandoffStages]; + array_aligned output_heads; + array_aligned output_head_checkpoints; + array_aligned, + Policy::kTmaWgmmaSmemAlignment> AG_bf16; + array_aligned, + Policy::kTmaWgmmaSmemAlignment> P_bf16; + array_aligned, + Policy::kTmaWgmmaSmemAlignment> W_bf16; + array_aligned, + Policy::kTmaWgmmaSmemAlignment> AGW_bf16; + array_aligned, + Policy::kBarrierAlignment> lower; + }; + + static constexpr size_t kSharedBytes = sizeof(SharedStorage); + static_assert(kSharedBytes == StateTraits::kPersistentSharedBytes); + // Commit reuses the unused query allocation for its final scaled-key operand. + static_assert(cosize_v >= cosize_v); + + template + struct NamedBarrier { + static_assert(static_cast(BarrierId) >= 0 + && static_cast(BarrierId) < kHardwareNamedBarrierCount); + static_assert(ParticipantWarps > 0); + + static CUTE_DEVICE void arrive() + { + asm volatile( + "barrier.arrive %0, %1;" + : + : "n"(static_cast(BarrierId)), "n"(ParticipantWarps * kWarpThreads) + : "memory"); + } + + static CUTE_DEVICE void sync() + { + asm volatile( + "barrier.sync %0, %1;" + : + : "n"(static_cast(BarrierId)), "n"(ParticipantWarps * kWarpThreads) + : "memory"); + } + }; + + template + struct StagedNamedBarrier { + static_assert(static_cast(FirstBarrierId) >= 0); + static_assert(static_cast(FirstBarrierId) + Stages <= kHardwareNamedBarrierCount); + static_assert(ParticipantWarps > 0); + + static CUTE_DEVICE void arrive(int stage) + { + const int barrier_id = static_cast(FirstBarrierId) + stage; + asm volatile( + "barrier.arrive %0, %1;" + : + : "r"(barrier_id), "n"(ParticipantWarps * kWarpThreads) + : "memory"); + } + + static CUTE_DEVICE void sync(int stage) + { + const int barrier_id = static_cast(FirstBarrierId) + stage; + asm volatile( + "barrier.sync %0, %1;" + : + : "r"(barrier_id), "n"(ParticipantWarps * kWarpThreads) + : "memory"); + } + }; + + using GramReady = + NamedBarrier; + using WReady = StagedNamedBarrier; + using AGReady = StagedNamedBarrier; + using StateDv0Ready = + NamedBarrier; + using StateDv1Ready = + NamedBarrier; + using PReady = + NamedBarrier; + using GatePrefixReady = + NamedBarrier; + using AGSolveReady = + NamedBarrier; + + template + static CUTE_DEVICE void ConvertAndReleaseStateDvHalf(Fp32StateTraits, + StateSource state_source, + StateOperand state_operand, + int state_conversion_thread, + StateDvHalf state_dv_half) + { + auto state_source_tile = + local_tile(state_source, + StateConversionTile{}, + make_coord(state_dv_half, _0{})); + auto state_operand_tile = + local_tile(state_operand, + StateConversionTile{}, + make_coord(state_dv_half, _0{})); + auto state_source_partition = + state_source_tile.compose(StateConversionThreadValueLayout{}); + auto state_operand_partition = + state_operand_tile.compose(StateConversionThreadValueLayout{}); + CUTE_UNROLL + for (int iteration = 0; iteration < StateConversionIterations::value; ++iteration) { + auto source_registers = make_tensor(Shape{}); + copy(Copy_Atom, float>{}, + state_source_partition(state_conversion_thread, make_coord(_, iteration)), + source_registers); + auto converted_registers = make_tensor(Shape{}); + recast>( + converted_registers)(0) = + cutlass::NumericArrayConverter{}( + recast>( + source_registers)(0)); + copy(Copy_Atom, Element>{}, + converted_registers, + state_operand_partition(state_conversion_thread, make_coord(_, iteration))); + } + cutlass::arch::fence_view_async_shared(); + // This WG releases the fenced Dv half to the peer WG without waiting for it. + StateReady::arrive(); + } + + template + static CUTE_DEVICE void ConvertAndReleaseStateDvHalf(Bf16StateTraits, + StateSource, + StateOperand, + int, + StateDvHalf) + { + } + + template + static CUTE_DEVICE void AcquireStateDvHalf(Fp32StateTraits) + { + // This WG acquires the peer WG's converted Dv half before its WGMMA reads it. + StateReady::sync(); + } + + template + static CUTE_DEVICE void AcquireStateDvHalf(Bf16StateTraits) + { + } + + template + static CUTE_DEVICE void mma_qk_by_k_tile(QKTensor qk, + KeyTensor key, + Accumulator& accumulator, + int lane, + int k_tile_index) + { + auto mma = QKTiledMma{}; + auto thread_mma = mma.get_thread_slice(lane); + auto key_tile = local_tile(key, QKMmaTileNK{}, make_coord(k_tile_index, _0{})); + // MMA A is [M=16,K=16]. K8 loads its valid [M=8,K=16] half with + // LDSM.x2; the cleared upper half supplies zero rows to mma.sync. + auto qk_mma_operand = make_tensor(qk.data(), SmemLayoutQKMmaOperandTile{}); + auto qk_fragment = thread_mma.partition_fragment_A(qk_mma_operand); + auto key_fragment = thread_mma.partition_fragment_B(key_tile); + auto copy_a = QKTiledCopyA{}; + auto copy_b = QKTiledCopyBKMajor{}; + auto thread_copy_a = copy_a.get_thread_slice(lane); + auto thread_copy_b = copy_b.get_thread_slice(lane); + auto qk_source = thread_copy_a.partition_S(qk); + auto key_source = thread_copy_b.partition_S(key_tile); + auto qk_valid_fragment = filter_zeros( + qk_fragment(QKMmaAFragmentMCoordinate{}, _, _)); + auto qk_source_by_k_block = group_modes<0, decltype(rank(qk_source))::value - 1>( + qk_source); + auto qk_fragment_by_k_block = + group_modes<0, decltype(rank(qk_valid_fragment))::value - 1>( + qk_valid_fragment); + auto key_view = thread_copy_b.retile_D(key_fragment); + clear(qk_fragment); + CUTE_UNROLL + // X[N,N] += QK[N,Dk] * K^T[Dk,N] + for (int block = 0; block < size<2>(qk_fragment); ++block) { + copy(QKCopyAtomA{}, + qk_source_by_k_block(_, block), + qk_fragment_by_k_block(_, block)); + copy(copy_b, key_source(_, _, block), key_view(_, _, block)); + cute::gemm(thread_mma, + qk_fragment(_, _, block), + key_fragment(_, _, block), + accumulator); + } + } + + template + static CUTE_DEVICE auto ComputeAGInverse(LowerTensor lower, int lane, int warp) + { + if (warp < AGBlocks::value && lane == 0) { + auto diagonal = local_tile(lower, AGBlockTile{}, make_coord(warp, warp)); + auto lower_block = make_tensor(AGBlockTranspose{}); + auto inverse = make_tensor(AGBlockTranspose{}); + CUTE_UNROLL + for (int row = 0; row < AGBlock::value; ++row) { + copy(Copy_Atom, float>{}, + diagonal(row, _), + lower_block(row, _)); + CUTE_UNROLL + for (int column = 0; column < AGBlock::value; ++column) { + inverse(row, column) = row == column ? 1.0f : 0.0f; + } + } + CUTE_UNROLL + for (int row = 1; row < AGBlock::value; ++row) { + CUTE_UNROLL + for (int column = 0; column < row; ++column) { + float value = 0.0f; + CUTE_UNROLL + for (int middle = 0; middle < row; ++middle) { + value -= lower_block(row, middle) * inverse(middle, column); + } + inverse(row, column) = value; + } + } + copy(inverse, diagonal); + } + // WG0 publishes the inverted 4x4 diagonal blocks. + AGSolveReady::sync(); + + // Merge adjacent 4x4 inverses into 8x8, then merge the two 8x8 + // inverses into the complete 16x16 inverse. + CUTE_UNROLL + for (int half_blocks = 1; half_blocks < AGBlocks::value; half_blocks *= _2::value) { + const int blocks_per_merge = _2::value * half_blocks; + const int merges = AGBlocks::value / blocks_per_merge; + auto merge_layout = make_layout( + make_shape(half_blocks, half_blocks, merges)); + + if (warp < size(merge_layout)) { + const auto merge_coordinate = merge_layout.get_hier_coord(warp); + const int row_in_half = int(get<0>(merge_coordinate)); + const int column_in_half = int(get<1>(merge_coordinate)); + const int merge = int(get<2>(merge_coordinate)); + const int first_block = merge * blocks_per_merge; + const int block_row = first_block + half_blocks + row_in_half; + const int block_column = first_block + column_in_half; + + if (lane < AGBlockThreads::value) { + auto mma = AGBlockMma{}; + auto thread_mma = mma.get_thread_slice(lane); + auto identity = make_identity_tensor( + make_shape(AGBlock{}, AGBlock{})); + auto coordinates = thread_mma.partition_C(identity); + auto right_product = thread_mma.make_fragment_C(coordinates); + clear(right_product); + CUTE_UNROLL + for (int inner = 0; inner < half_blocks; ++inner) { + const int middle_block = first_block + half_blocks + inner; + auto right_inverse = local_tile( + lower, AGBlockTile{}, make_coord(block_row, middle_block)); + auto cross_block = local_tile( + lower, AGBlockTile{}, make_coord(middle_block, block_column)); + auto cross_block_transpose = cross_block.compose(AGBlockTranspose{}); + cute::gemm(mma, + thread_mma.partition_A(right_inverse), + thread_mma.partition_B(cross_block_transpose), + right_product); + } + auto right_product_scratch = local_tile( + lower, AGBlockTile{}, make_coord(block_column, block_row)); + copy(right_product, thread_mma.partition_C(right_product_scratch)); + } + } + // The right-half products are stored in the unused upper blocks. + AGSolveReady::sync(); + + if (warp < size(merge_layout)) { + const auto merge_coordinate = merge_layout.get_hier_coord(warp); + const int row_in_half = int(get<0>(merge_coordinate)); + const int column_in_half = int(get<1>(merge_coordinate)); + const int merge = int(get<2>(merge_coordinate)); + const int first_block = merge * blocks_per_merge; + const int block_row = first_block + half_blocks + row_in_half; + const int block_column = first_block + column_in_half; + + if (lane < AGBlockThreads::value) { + auto mma = AGBlockMma{}; + auto thread_mma = mma.get_thread_slice(lane); + auto identity = make_identity_tensor( + make_shape(AGBlock{}, AGBlock{})); + auto coordinates = thread_mma.partition_C(identity); + auto inverse = thread_mma.make_fragment_C(coordinates); + clear(inverse); + CUTE_UNROLL + for (int inner = 0; inner < half_blocks; ++inner) { + const int middle_block = first_block + inner; + auto right_product = local_tile( + lower, AGBlockTile{}, make_coord(middle_block, block_row)); + auto left_inverse = local_tile( + lower, AGBlockTile{}, make_coord(middle_block, block_column)); + auto left_inverse_transpose = left_inverse.compose(AGBlockTranspose{}); + cute::gemm(mma, + thread_mma.partition_A(right_product), + thread_mma.partition_B(left_inverse_transpose), + inverse); + } + auto inverse_values = coalesce(inverse); + CUTE_UNROLL + for (int index = 0; index < size(inverse_values); ++index) { + inverse_values(index) = -inverse_values(index); + } + auto inverse_block = local_tile( + lower, AGBlockTile{}, make_coord(block_row, block_column)); + copy(inverse, thread_mma.partition_C(inverse_block)); + } + } + // All warps must finish reading the upper-block scratch before it is cleared. + AGSolveReady::sync(); + if (warp < size(merge_layout)) { + const auto merge_coordinate = merge_layout.get_hier_coord(warp); + const int row_in_half = int(get<0>(merge_coordinate)); + const int column_in_half = int(get<1>(merge_coordinate)); + const int merge = int(get<2>(merge_coordinate)); + const int first_block = merge * blocks_per_merge; + const int block_row = first_block + half_blocks + row_in_half; + const int block_column = first_block + column_in_half; + + if (lane < AGBlockThreads::value) { + auto thread_mma = AGBlockMma{}.get_thread_slice(lane); + auto right_product_scratch = local_tile( + lower, AGBlockTile{}, make_coord(block_column, block_row)); + clear(thread_mma.partition_C(right_product_scratch)); + } + } + // Publish a zero upper triangle for the next merge and final AG scaling. + AGSolveReady::sync(); + } + return lower; + } + + struct HeadCoordinate { + int request; + int value_head; + int query_head; + int head_group; + int local_head; + }; + + static CUTE_DEVICE HeadCoordinate DecodeHead(int flattened_head, + int batch, + int hq, + int hv, + int value_heads_per_query_head, + int num_head_groups, + int heads_per_block) + { + auto head_layout = make_layout(make_shape(hv, batch)); + const auto head_coordinate = head_layout.get_hier_coord(flattened_head); + const int value_head = int(get<0>(head_coordinate)); + const int request = int(get<1>(head_coordinate)); + + auto query_head_layout = make_layout(make_shape(value_heads_per_query_head, hq)); + const auto query_head_coordinate = query_head_layout.get_hier_coord(value_head); + + auto state_head_layout = make_layout(make_shape(heads_per_block, num_head_groups)); + const auto state_head_coordinate = state_head_layout.get_hier_coord(value_head); + return {request, + value_head, + int(get<1>(query_head_coordinate)), + int(get<1>(state_head_coordinate)), + int(get<0>(state_head_coordinate))}; + } + + template + static CUTE_DEVICE void StoreOutput(OutputCopy output_copy, + OutputFragment& output_fragment, + OutputCoordinates output_coordinates, + OutputTile output_tile, + StoreVectorElements, + ValidCoordinate, + int valid_positions, + int thread) + { + auto output_thread = output_copy.get_thread_slice(thread); + auto output_source = output_thread.retile_S(output_fragment); + auto output_destination = output_thread.partition_D(output_tile); + auto output_coordinate = output_thread.retile_S(output_coordinates); + auto output_source_vectors = + zipped_divide(output_source, make_tile(StoreVectorElements{})); + auto output_destination_vectors = + zipped_divide(output_destination, make_tile(StoreVectorElements{})); + auto output_coordinate_vectors = + zipped_divide(output_coordinate, make_tile(StoreVectorElements{})); + CUTE_UNROLL + for (int vector = 0; vector < size<1>(output_coordinate_vectors); ++vector) { + const auto coordinate = output_coordinate_vectors(_0{}, vector); + if (int(get(coordinate)) < valid_positions) { + auto output_vector = make_tensor(shape<0>(output_source_vectors)); + copy(output_source_vectors(_, vector), output_vector); + copy(output_copy, + output_vector, + output_destination_vectors(_, vector)); + } + } + } + + static CUTE_DEVICE void ProcessGramPAGHead(SharedStorage& shared_storage, + int read_stage, + int handoff_stage, + QKPipeline& qk_pipeline, + QKPipelineState qk_pipe_release, + int lane, + int warp, + int handoff_phase) + { + const int valid_positions = + kCommit ? shared_storage.output_heads[read_stage].commit_length : Capacity; + auto key_smem = make_tensor( + make_smem_ptr(shared_storage.k.data()), SmemLayoutQK{}); + auto key = key_smem(_, _, read_stage); + auto gate_exp_difference_storage = make_tensor( + make_smem_ptr(shared_storage.gate_exp_difference.data()), + SmemLayoutGateExpDifference{}); + auto gate_exp_difference = gate_exp_difference_storage(_, _, handoff_stage); + + auto AG_storage = make_tensor( + make_smem_ptr(shared_storage.AG_bf16.data()), SmemLayoutAG{}); + auto AG = AG_storage(_, _, handoff_stage); + auto beta_smem = make_tensor( + make_smem_ptr(shared_storage.beta.data()), SmemLayoutBeta{}); + auto beta = beta_smem(_, _0{}, read_stage); + auto beta_checkpoint_storage = make_tensor( + make_smem_ptr(shared_storage.beta_checkpoint.data()), + SmemLayoutBetaCheckpoint{}); + auto beta_checkpoint = beta_checkpoint_storage(_, read_stage); + auto lower_storage = make_tensor( + make_smem_ptr(shared_storage.lower.data()), SmemLayoutLowerTriangular{}); + auto lower = lower_storage(_, _, read_stage); + + if (warp == kAGWarp && lane < kMmaM) { + beta_checkpoint(lane) = beta(lane); + } + + // _gram_: Gram[N,N] = diag(beta) * StrictLower(K[N,Dk] * K^T[Dk,N]) + const int gram_tile = warp - kGramWarpBegin; + if (gram_tile < kCapacityNTiles) { + const int col0 = gram_tile * kMmaN; + auto mma = QKTiledMma{}; + auto thr = mma.get_thread_slice(lane); + auto identity = make_identity_tensor(make_shape(MmaM{}, MmaN{})); + auto coords = thr.partition_C(identity); + auto rC = thr.make_fragment_C(coords); + clear(rC); + mma_qk_by_k_tile(key, key, rC, lane, gram_tile); + auto values = coalesce(rC); + auto coordinates = coalesce(coords); + CUTE_UNROLL + for (int index = 0; index < size(values); ++index) { + const auto coordinate = coordinates(index); + const int row = int(get<0>(coordinate)); + const int column = col0 + int(get<1>(coordinate)); + values(index) = column < row && (!kCommit || row < valid_positions) ? beta(row) * values(index) : 0.0f; + } + auto lower_tiled_copy = make_tiled_copy_C( + Copy_Atom, float>{}, mma); + auto lower_thread_copy = lower_tiled_copy.get_thread_slice(lane); + auto lower_source = lower_thread_copy.retile_S(rC); + auto lower_destination = lower_thread_copy.partition_D(lower); + copy(lower_tiled_copy, + lower_source(_, _0{}, _0{}), + lower_destination(_, _0{}, gram_tile)); + } + + // The Gram-role warps publish the complete Gram matrix before the AG solver starts. + GramReady::sync(); + // The Gram/AG warps release their Q/K-stage participation after Gram. + qk_pipeline.consumer_release(qk_pipe_release); + + auto AG_inverse = ComputeAGInverse(lower, lane, warp); + // _ag_: WG0 owns the final AG scaling. + // WG0 acquires the completed prefix before _ag_scale_. + wait_barrier(shared_storage.prefix_ready[handoff_stage], handoff_phase); + + // _ag_scale_: AG[row,column] *= exp(prefix[row] - prefix[column]) * beta[column] + if (lane < size(AGScaleLaneCoordinates{})) { + const auto coordinate = AGScaleLaneCoordinates{}.get_hier_coord(lane); + const int row = int(AGScaleRowCoordinates{}(warp, get<0>(coordinate))); + const int column_vector = int(get<1>(coordinate)); + auto inverse = make_tensor(Shape{}); + auto gate_difference = make_tensor(Shape{}); + auto beta_vector = make_tensor(Shape{}); + copy(AGScaleLoadAtom{}, + coalesce(local_tile( + AG_inverse, AGScaleTile{}, make_coord(row, column_vector))), + inverse); + copy(AGScaleLoadAtom{}, + coalesce(local_tile( + gate_exp_difference, + AGScaleTile{}, + make_coord(row, column_vector))), + gate_difference); + copy(AGScaleLoadAtom{}, + local_tile(beta_checkpoint, + Tile{}, + make_coord(column_vector)), + beta_vector); + auto transformed = make_tensor(Shape{}); + CUTE_UNROLL + for (int element = 0; element < AGScaleVectorElements::value; ++element) { + const int column = column_vector * AGScaleVectorElements::value + element; + transformed(element) = kCommit && (row >= valid_positions || column > row) ? Element{} : + static_cast(inverse(element) * gate_difference(element) * beta_vector(element)); + } + copy(AGScaleStoreAtom{}, + transformed, + coalesce(local_tile( + AG, AGScaleTile{}, make_coord(row, column_vector)))); + } + // WG0 releases its complete AG shared-memory tile to WG1. + AGReady::arrive(handoff_stage); + } + + static CUTE_DEVICE void ProcessFinalHead(SharedStorage& shared_storage, + int read_stage, + int handoff_stage, + QKPipeline& qk_pipeline, + QKPipelineState qk_pipe_release, + SVPipeline& sv_pipeline, + SVPipelineState sv_pipe_read, + SVPipelineState sv_pipe_release, + __nv_bfloat16* out, + int valid_positions, + int batch, + int hv, + int64_t out_batch_stride, + int64_t out_token_stride, + int64_t out_head_stride, + int lane, + int warp, + int warp_group_thread, + State* commit_state, + int commit_length) + { + auto query_smem = make_tensor( + make_smem_ptr(shared_storage.q.data()), SmemLayoutQK{}); + auto key_smem = make_tensor( + make_smem_ptr(shared_storage.k.data()), SmemLayoutQK{}); + auto query = query_smem(_, _, read_stage); + auto key = key_smem(_, _, read_stage); + + // _p_: P[N,N] = Q[N,Dk] * K^T[Dk,N] + const int P_tile = warp_group_thread / kWarpThreads; + auto P_mma = QKTiledMma{}; + auto P_thread = P_mma.get_thread_slice(lane); + auto P_identity = make_identity_tensor(make_shape(MmaM{}, MmaN{})); + auto P_coordinates = P_thread.partition_C(P_identity); + auto P_accumulator = P_thread.make_fragment_C(P_coordinates); + clear(P_accumulator); + if (!kCommit && P_tile < kCapacityNTiles) { + mma_qk_by_k_tile(query, key, P_accumulator, lane, P_tile); + } + + if (!kCommit && warp == kOutputHeadCheckpointWarp && lane == 0) { + shared_storage.output_head_checkpoints[read_stage] = + shared_storage.output_heads[read_stage]; + } + auto output_query = inner_partition(query, Tile{}, _0{}); + auto gate_exp_storage = make_tensor( + make_smem_ptr(shared_storage.gate_exp.data()), SmemLayoutGateExp{}); + auto gate_exp = gate_exp_storage(_, handoff_stage); + auto gate_exp_difference_storage = make_tensor( + make_smem_ptr(shared_storage.gate_exp_difference.data()), + SmemLayoutGateExpDifference{}); + auto gate_exp_difference = gate_exp_difference_storage(_, _, handoff_stage); + + // WG1 acquires state and value before converting the state operand. + sv_pipeline.consumer_wait(sv_pipe_read); + + auto mma = SSWgmmaTiledMma{}; + auto thread_mma = mma.get_thread_slice(warp_group_thread); + auto output_identity = make_identity_tensor(make_shape(BlockDv{}, OutputN{})); + auto output_coordinates = thread_mma.partition_C(output_identity); + auto output = thread_mma.make_fragment_C(output_coordinates); + clear(output); + + auto state = make_tensor( + make_smem_ptr(reinterpret_cast(shared_storage.state.data())), + SmemLayoutState{}); + auto state_source = state(_, _, _, _, read_stage).compose(StateMmaLayout{}); + auto state_operand = make_tensor( + make_smem_ptr(reinterpret_cast( + shared_storage.state.data() + StateTraits::kStateOperandOffsetBytes + + read_stage * StateTraits::kStateOperandStageBytes)), + SmemLayoutStateOperand{}); + auto query_fragment = thread_mma.make_fragment_B( + thread_mma.partition_B(output_query)); + auto output_dv_halves = zipped_divide( + output, make_tile(shape<0>(output), _1{}, shape<2>(output))); + + // _o_sq_: O0[Dv,N] = S[Dv,Dk] * Q^T[Dk,N] + if constexpr (!kCommit) { + warpgroup_fence_operand(output); + } + auto mma_state_dv_half = [&](auto state_dv_half) { + auto state_tile = local_tile( + state_operand, StateConversionTile{}, make_coord(state_dv_half, _0{})); + auto state_fragment = thread_mma.make_fragment_A( + thread_mma.partition_A(state_tile)); + auto output_half = output_dv_halves( + make_coord(_, _, _), make_coord(_0{}, state_dv_half, _0{})); + warpgroup_arrive(); + cute::gemm(mma, state_fragment, query_fragment, output_half); + warpgroup_commit_batch(); + }; + // _state_conversion_: WG1 converts S'[Dv=0:64,Dk]. + ConvertAndReleaseStateDvHalf(StateTraits{}, + state_source, + state_operand, + warp_group_thread, + _0{}); + if constexpr (!kCommit) { + mma_state_dv_half(_0{}); + } + AcquireStateDvHalf(StateTraits{}); + if constexpr (!kCommit) { + mma_state_dv_half(_1{}); + warpgroup_wait<0>(); + warpgroup_fence_operand(output); + // Verification releases its inputs after the state/query product. + qk_pipeline.consumer_release(qk_pipe_release); + sv_pipeline.consumer_release(sv_pipe_release); + } + + // WG1 acquires WG2's W tile and WG0's AG tile before _agw_. + WReady::sync(handoff_stage); + AGReady::sync(handoff_stage); + auto W_storage = make_tensor( + make_smem_ptr(shared_storage.W_bf16.data()), SmemLayoutW{}); + auto AG_storage = make_tensor( + make_smem_ptr(shared_storage.AG_bf16.data()), SmemLayoutAG{}); + auto W_smem = W_storage(_, _, handoff_stage); + auto AG = AG_storage(_, _, handoff_stage); + auto output_AG = inner_partition(AG, Tile{}, _0{}); + auto W_fragment = thread_mma.make_fragment_A( + thread_mma.partition_A(W_smem)); + auto AG_fragment = thread_mma.make_fragment_B( + thread_mma.partition_B(output_AG)); + auto AGW = thread_mma.make_fragment_C(output_coordinates); + clear(AGW); + + // _agw_: AGW[Dv,N] = W[Dv,N] * AG[N,N] + cutlass::arch::fence_view_async_shared(); + warpgroup_fence_operand(AGW); + warpgroup_arrive(); + cute::gemm(mma, W_fragment, AG_fragment, AGW); + warpgroup_commit_batch(); + + if constexpr (!kCommit) { + // _o_scale_: O1[Dv,N] = scale*exp(gate) * O0[Dv,N] + auto output_values = coalesce(output); + auto output_coords = coalesce(output_coordinates); + CUTE_UNROLL + for (int index = 0; index < size(output_values); ++index) { + const int row = int(get<1>(output_coords(index))); + const float scale = kHeadScale * gate_exp(row); + output_values(index) *= scale; + } + + auto P_storage = make_tensor( + make_smem_ptr(shared_storage.P_bf16.data()), SmemLayoutP{}); + auto P = P_storage; + // _p_scale_: P'[row,column] = scale * P[row,column] * exp(prefix[row] - prefix[column]) + if (P_tile < kCapacityNTiles) { + auto P_packed = make_fragment_like(P_accumulator); + auto P_values = coalesce( + P_accumulator(TokenRows{}, _, _)); + auto P_packed_values = coalesce( + P_packed(TokenRows{}, _, _)); + auto P_coords = coalesce( + P_coordinates(TokenRows{}, _, _)); + auto P_value_pairs = zipped_divide(P_values, Tile<_2>{}); + auto P_packed_pairs = zipped_divide(P_packed_values, Tile<_2>{}); + auto P_coordinate_pairs = zipped_divide(P_coords, Tile<_2>{}); + CUTE_UNROLL + for (int pair = 0; pair < size<1>(P_value_pairs); ++pair) { + const auto coordinate = P_coordinate_pairs(_0{}, pair); + const int row = int(get<0>(coordinate)); + const int column = P_tile * kMmaN + int(get<1>(coordinate)); + auto gate_difference = make_tensor(Shape<_2>{}); + copy(Copy_Atom, float>{}, + coalesce(local_tile(gate_exp_difference, + Tile<_1, _2>{}, + make_coord(row, column / _2::value))), + gate_difference); + CUTE_UNROLL + for (int element = 0; element < size<0>(P_value_pairs); ++element) { + const float P_value = kHeadScale * P_value_pairs(element, pair) + * gate_difference(element); + P_packed_pairs(element, pair) = static_cast(P_value); + } + } + auto P_tiled_copy = make_tiled_copy_C( + Copy_Atom{}, P_mma); + auto P_thread_copy = P_tiled_copy.get_thread_slice(lane); + auto P_destination = P_thread_copy.partition_D( + as_position_independent_swizzle_tensor(P)); + auto P_source = P_thread_copy.retile_S(P_packed); + copy(P_tiled_copy, + P_source(_, _0{}, _0{}), + P_destination(_, _0{}, P_tile)); + } + } + warpgroup_wait<0>(); + warpgroup_fence_operand(AGW); + if constexpr (!kCommit) { + // Every WG1 warp acquires the complete scaled P matrix. + PReady::sync(); + } + + auto AGW_packed = make_fragment_like(AGW); + CUTE_UNROLL + for (int index = 0; index < size(AGW); ++index) { + AGW_packed(index) = static_cast(AGW(index)); + } + auto AGW_storage = make_tensor( + make_smem_ptr(shared_storage.AGW_bf16.data()), SmemLayoutAGW{}); + auto AGW_smem = AGW_storage; + auto output_AGW_smem = inner_partition( + AGW_smem, Tile{}, _0{}); + auto AGW_tiled_copy = make_tiled_copy_C( + WgmmaAccumulatorStoreAtom{}, mma); + auto AGW_thread_copy = AGW_tiled_copy.get_thread_slice(warp_group_thread); + auto AGW_destination = AGW_thread_copy.partition_D( + as_position_independent_swizzle_tensor(output_AGW_smem)); + auto AGW_source = AGW_thread_copy.retile_S(AGW_packed); + copy(AGW_tiled_copy, AGW_source, AGW_destination); + + if constexpr (kCommit) { + // AGW[Dv,t] is the solved delta for transition t. The final state is + // exp(prefix[L-1])*S0 + AGW * (exp(prefix[L-1]-prefix[t])*K[t,Dk]). + // Both update operands have reduction extent 16, including K8's zero tail. + auto commit_key = make_tensor( + make_smem_ptr(shared_storage.q.data()), SmemLayoutCommitKey{}); + auto key_transpose = make_tensor(key.data(), select<1, 0>(key.layout())); + auto key_copy = CommitKeyTiledCopy{}; + auto key_thread_copy = key_copy.get_thread_slice(warp_group_thread); + auto key_source = key_thread_copy.partition_S(key_transpose); + auto key_destination = key_thread_copy.partition_D(commit_key); + auto key_coordinates = key_thread_copy.partition_D( + make_identity_tensor(Shape{})); + // Each lane owns eight adjacent Dk elements. The SW128 layout + // permutes the eight 16-byte vectors within each 64-element box. + CUTE_UNROLL + for (int tile = 0; tile < size<2>(key_destination); ++tile) { + const int position = int(get<1>(key_coordinates(_0{}, _0{}, tile))); + auto values = make_fragment_like(key_destination(_, _0{}, tile)); + clear(values); + if (position < commit_length) { + copy(CommitKeyCopyAtom{}, key_source(_, _0{}, tile), values); + const float decay = gate_exp_difference(commit_length - 1, position); + CUTE_UNROLL + for (int index = 0; index < size(values); ++index) { + values(index) = static_cast(float(values(index)) * decay); + } + } + copy(CommitKeyCopyAtom{}, values, key_destination(_, _0{}, tile)); + } + // The final WG owns both scratch operands through the state WGMMA wait. + PReady::sync(); + cutlass::arch::fence_view_async_shared(); + + auto update_mma = StateUpdateTiledMma{}; + auto update_thread = update_mma.get_thread_slice(warp_group_thread); + auto state_destination = make_tensor( + make_gmem_ptr(commit_state), Layout, Stride<_1, BlockDv>>{}); + const float state_decay = gate_exp(commit_length - 1); + // One 64x64 tile keeps only 32 FP32 accumulator values live per lane. + CUTLASS_PRAGMA_NO_UNROLL + for (int tile = 0; tile < size(StateUpdateTiles{}); ++tile) { + const auto coordinate = StateUpdateTiles{}.get_hier_coord(tile); + auto state_tile = local_tile(state_source, StateUpdateTile{}, coordinate); + auto state_input = update_thread.partition_C(state_tile); + auto accumulator = update_thread.make_fragment_C(state_input); + CUTE_UNROLL + for (int index = 0; index < size(accumulator); ++index) { + // Keep the original FP32 state for the residual; round BF16 only at the final store. + accumulator(index) = state_decay * float(state_input(index)); + } + auto delta_tile = local_tile( + AGW_smem, Tile{}, make_coord(get<0>(coordinate), _0{})); + auto key_tile = local_tile( + commit_key, Tile{}, make_coord(get<1>(coordinate), _0{})); + auto delta_fragment = update_thread.make_fragment_A(update_thread.partition_A(delta_tile)); + auto update_key = update_thread.make_fragment_B(update_thread.partition_B(key_tile)); + warpgroup_fence_operand(accumulator); + warpgroup_arrive(); + cute::gemm(update_mma, delta_fragment, update_key, accumulator); + warpgroup_commit_batch(); + warpgroup_wait<0>(); + warpgroup_fence_operand(accumulator); + + auto destination_tile = local_tile( + state_destination, StateUpdateTile{}, coordinate); + auto destination = update_thread.partition_C(destination_tile); + CUTE_UNROLL + for (int index = 0; index < size(accumulator); ++index) { + destination(index) = static_cast(accumulator(index)); + } + } + // Commit retains K and the unconverted entry state until the final update completes. + qk_pipeline.consumer_release(qk_pipe_release); + sv_pipeline.consumer_release(sv_pipe_release); + } + + if constexpr (!kCommit) { + auto P = make_tensor(make_smem_ptr(shared_storage.P_bf16.data()), SmemLayoutP{}); + auto output_P = inner_partition(P, Tile{}, _0{}); + auto AGW_fragment = thread_mma.make_fragment_A( + thread_mma.partition_A(AGW_smem)); + auto P_fragment = thread_mma.make_fragment_B( + thread_mma.partition_B(output_P)); + + // _o_final_: O2[Dv,N] = O1[Dv,N] + AGW[Dv,N] * P[N,N] + cutlass::arch::fence_view_async_shared(); + warpgroup_fence_operand(output); + warpgroup_arrive(); + cute::gemm(mma, AGW_fragment, P_fragment, output); + warpgroup_commit_batch(); + warpgroup_wait<0>(); + warpgroup_fence_operand(output); + + auto output_tensor = make_tensor( + make_gmem_ptr(reinterpret_cast(out)), + make_layout(make_shape(batch, hv, BlockDv{}, OutputN{}), + make_stride(out_batch_stride, + out_head_stride, + _1{}, + out_token_stride))); + const auto output_head = shared_storage.output_head_checkpoints[read_stage]; + auto output_tile = output_tensor(output_head.request, output_head.value_head, _, _); + StoreOutput(OutputStoreTiledCopy{}, + output, + output_coordinates, + output_tile, + OutputStoreVectorElements{}, + _1{}, + valid_positions, + warp_group_thread); + } + } + + static CUTE_DEVICE void ProcessUpdateHead(SharedStorage& shared_storage, + int read_stage, + int handoff_stage, + QKPipeline& qk_pipeline, + QKPipelineState qk_pipe_release, + SVPipeline& sv_pipeline, + SVPipelineState sv_pipe_read, + SVPipelineState sv_pipe_release, + int lane, + int warp, + int warp_group_thread, + int handoff_phase) + { + const int commit_length = kCommit ? shared_storage.output_heads[read_stage].commit_length : Capacity; + auto key_smem = make_tensor( + make_smem_ptr(shared_storage.k.data()), SmemLayoutQK{}); + auto key = key_smem(_, _, read_stage); + auto output_key = inner_partition(key, Tile{}, _0{}); + auto gate_smem = make_tensor( + make_smem_ptr(shared_storage.gate.data()), SmemLayoutGate{}); + auto gate = gate_smem(_, _0{}, read_stage); + auto gate_exp_storage = make_tensor( + make_smem_ptr(shared_storage.gate_exp.data()), SmemLayoutGateExp{}); + auto gate_exp = gate_exp_storage(_, handoff_stage); + auto gate_exp_difference_storage = make_tensor( + make_smem_ptr(shared_storage.gate_exp_difference.data()), + SmemLayoutGateExpDifference{}); + auto gate_exp_difference = gate_exp_difference_storage(_, _, handoff_stage); + + // _prefix_: Prefix sum of gate. + auto gate_prefix = make_tensor(gate.data(), Layout>{}); + if (warp == kPrefixWarp) { + float prefix = lane < Capacity ? gate_prefix(lane) : 0.0f; + CUTE_UNROLL + for (int delta = 1; delta < Capacity; delta <<= 1) { + const float other = __shfl_up_sync(kWarpMask, prefix, delta); + if (lane >= delta) { + prefix += other; + } + } + if (lane < Capacity) { + gate_prefix(lane) = prefix; + gate_exp(lane) = FastExp(prefix); + } + } + + // WG2 acquires the stored prefix vector before cooperatively materializing its differences. + GatePrefixReady::sync(); + const auto prefix_thread_coordinate = + PrefixWarpGroupThreadLayout{}.get_hier_coord(warp_group_thread); + const int prefix_warp = int(get<0>(prefix_thread_coordinate)); + if (lane < size(PrefixOutputLaneLayout{})) { + const auto prefix_output_coordinate = + PrefixOutputLaneLayout{}.get_hier_coord(lane); + const int row = int(PrefixOutputRowLayout{}( + prefix_warp, get<0>(prefix_output_coordinate))); + const int column_vector = int(get<1>(prefix_output_coordinate)); + auto prefix_columns = make_tensor(Shape<_4>{}); + copy(Copy_Atom, float>{}, + local_tile(gate_prefix, Tile<_4>{}, make_coord(column_vector)), + prefix_columns); + const float prefix_row = gate_prefix(row); + auto gate_exp_difference_vector = make_tensor(Shape<_4>{}); + CUTE_UNROLL + for (int element = 0; element < size(gate_exp_difference_vector); ++element) { + const int column = int(PrefixColumnVectorLayout{}(column_vector, element)); + gate_exp_difference_vector(element) = + column <= row ? FastExp(prefix_row - prefix_columns(element)) : 0.0f; + } + auto gate_exp_difference_destination = coalesce(local_tile( + gate_exp_difference, + Tile<_1, _4>{}, + make_coord(row, column_vector))); + copy(Copy_Atom, float>{}, + gate_exp_difference_vector, + gate_exp_difference_destination); + } + // Every WG2 warp releases its completed columns to the prefix consumers. + arrive_barrier(shared_storage.prefix_ready[handoff_stage]); + + // WG2 acquires state and value after completing the gate-only prefix phase. + sv_pipeline.consumer_wait(sv_pipe_read); + + auto state_mma = SSWgmmaTiledMma{}; + auto state_thread_mma = state_mma.get_thread_slice(warp_group_thread); + auto output_identity = make_identity_tensor(make_shape(BlockDv{}, OutputN{})); + auto output_coordinates = state_thread_mma.partition_C(output_identity); + auto U = state_thread_mma.make_fragment_C(output_coordinates); + clear(U); + + auto state = make_tensor( + make_smem_ptr(reinterpret_cast(shared_storage.state.data())), + SmemLayoutState{}); + auto state_source = state(_, _, _, _, read_stage).compose(StateMmaLayout{}); + auto state_operand = make_tensor( + make_smem_ptr(reinterpret_cast( + shared_storage.state.data() + StateTraits::kStateOperandOffsetBytes + + read_stage * StateTraits::kStateOperandStageBytes)), + SmemLayoutStateOperand{}); + auto key_fragment = state_thread_mma.make_fragment_B( + state_thread_mma.partition_B(output_key)); + auto U_dv_halves = zipped_divide( + U, make_tile(shape<0>(U), _1{}, shape<2>(U))); + + // _u_: U[Dv,N] = S[Dv,Dk] * K^T[Dk,N] + warpgroup_fence_operand(U); + auto mma_state_dv_half = [&](auto state_dv_half) { + auto state_tile = local_tile( + state_operand, StateConversionTile{}, make_coord(state_dv_half, _0{})); + auto state_fragment = state_thread_mma.make_fragment_A( + state_thread_mma.partition_A(state_tile)); + auto U_half = U_dv_halves( + make_coord(_, _, _), make_coord(_0{}, state_dv_half, _0{})); + warpgroup_arrive(); + cute::gemm(state_mma, state_fragment, key_fragment, U_half); + warpgroup_commit_batch(); + }; + // _state_conversion_: WG2 converts S'[Dv=64:128,Dk]. + ConvertAndReleaseStateDvHalf(StateTraits{}, + state_source, + state_operand, + warp_group_thread, + _1{}); + mma_state_dv_half(_1{}); + AcquireStateDvHalf(StateTraits{}); + mma_state_dv_half(_0{}); + warpgroup_wait<0>(); + warpgroup_fence_operand(U); + + // WG2 releases Q/K after the final key WGMMA read. + qk_pipeline.consumer_release(qk_pipe_release); + + // WG2 independently acquires the completed prefix before _w_. + wait_barrier(shared_storage.prefix_ready[handoff_stage], handoff_phase); + + // _w_: W[Dv,N] = V^T[Dv,N] - exp(g)*U[Dv,N] + auto W = make_fragment_like(U); + auto W_values = coalesce(W); + auto U_values = coalesce(U); + auto output_coords = coalesce(output_coordinates); + auto input_values = make_tensor( + make_smem_ptr(shared_storage.value.data()), SmemLayoutValue{}); + CUTE_UNROLL + for (int index = 0; index < size(W_values); ++index) { + const auto coordinate = output_coords(index); + const int dv = int(get<0>(coordinate)); + const int row = int(get<1>(coordinate)); + const float gate_scale = gate_exp(row); + const auto dv_coordinate = ValueDvCoordinateLayout{}.get_hier_coord(dv); + W_values(index) = kCommit && row >= commit_length ? Element{} : static_cast( + float(input_values(get<0>(dv_coordinate), + row, + get<1>(dv_coordinate), + read_stage)) + - gate_scale * U_values(index)); + } + // WG2 releases state/value after its final value read. + sv_pipeline.consumer_release(sv_pipe_release); + + auto W_storage = make_tensor( + make_smem_ptr(shared_storage.W_bf16.data()), SmemLayoutW{}); + auto W_smem = W_storage(_, _, handoff_stage); + auto output_W_smem = inner_partition(W_smem, Tile{}, _0{}); + auto W_tiled_copy = make_tiled_copy_C( + WgmmaAccumulatorStoreAtom{}, state_thread_mma); + auto W_thread_copy = W_tiled_copy.get_thread_slice(warp_group_thread); + auto W_destination = W_thread_copy.partition_D( + as_position_independent_swizzle_tensor(output_W_smem)); + auto W_source = W_thread_copy.retile_S(W); + copy(W_tiled_copy, W_source, W_destination); + // WG2 releases its complete W shared-memory tile to WG1. + WReady::arrive(handoff_stage); + } + +public: + template + struct DeviceOperator { + using SharedStorage = typename Policy::SharedStorage; + + static constexpr int MaxThreadsPerBlock = Policy::kThreads; + static constexpr int MinBlocksPerMultiprocessor = 1; + static constexpr int SharedMemoryAlignment = alignof(SharedStorage); + + struct KernelParams { + const float* g; + const float* beta; + const TmaDescriptor* state_tma_descs; + __nv_bfloat16* out; + int valid_positions; + int batch; + int hq; + int hv; + int value_heads_per_query_head; + int num_head_groups; + int heads_per_block; + int64_t g_batch_stride; + int64_t g_token_stride; + int64_t beta_batch_stride; + int64_t beta_token_stride; + int64_t out_batch_stride; + int64_t out_token_stride; + int64_t out_head_stride; + int state_layer; + void* const* state_ptrs; + const int* commit_lengths; + const bool* finished; + int64_t state_request_stride; + int64_t state_group_stride; + }; + + struct Params { + TmaQ tma_q; + TmaK tma_k; + TmaV tma_v; + TmaState tma_state; + KernelParams kernel; + }; + + CUTE_DEVICE void operator()(const Params& parameters_, SharedStorage& storage) + { + const int thread = static_cast(threadIdx.x); + const int lane = thread % Policy::kWarpThreads; + const int warp = thread / Policy::kWarpThreads; + const int warp_group = thread / cutlass::NumThreadsPerWarpGroup; + const int warp_group_thread = thread % cutlass::NumThreadsPerWarpGroup; + const bool is_producer_group = + warp_group == static_cast(WarpGroupRole::Producer); + const bool is_producer_warp = is_producer_group && warp == kProducerWarp; + const bool is_sv_consumer = + warp_group == static_cast(WarpGroupRole::Output) + || warp_group == static_cast(WarpGroupRole::Update); + + const auto& parameters = parameters_.kernel; + typename QKPipeline::Params qk_pipeline_params; + qk_pipeline_params.transaction_bytes = kQKTransactionBytes; + qk_pipeline_params.role = is_producer_warp ? QKPipeline::ThreadCategory::Producer : + is_producer_group ? QKPipeline::ThreadCategory::NonParticipant : + QKPipeline::ThreadCategory::Consumer; + qk_pipeline_params.is_leader = thread == kProducerThread; + qk_pipeline_params.num_consumers = kComputeThreads; + qk_pipeline_params.num_producers = 1 + kWarpThreads; + qk_pipeline_params.initializing_warp = 0; + QKPipeline qk_pipeline( + storage.qk_pipeline, + qk_pipeline_params, + PipelineShape{}, + bool_constant{}, + bool_constant{}); + + typename SVPipeline::Params sv_pipeline_params; + sv_pipeline_params.transaction_bytes = kSVTransactionBytes; + sv_pipeline_params.role = is_producer_warp ? SVPipeline::ThreadCategory::Producer : + is_sv_consumer ? SVPipeline::ThreadCategory::Consumer : + SVPipeline::ThreadCategory::NonParticipant; + sv_pipeline_params.is_leader = thread == kProducerThread; + sv_pipeline_params.num_consumers = kSVConsumerThreads; + sv_pipeline_params.num_producers = 1; + sv_pipeline_params.initializing_warp = 0; + SVPipeline sv_pipeline( + storage.sv_pipeline, + sv_pipeline_params, + PipelineShape{}, + bool_constant{}, + bool_constant{}); + + if (thread == 0) { + CUTE_UNROLL + for (int stage = 0; stage < kStages; ++stage) { + storage.output_head_ready[stage].init(1); + } + CUTE_UNROLL + for (int stage = 0; stage < kHandoffStages; ++stage) { + initialize_barrier(storage.prefix_ready[stage], cutlass::NumThreadsPerWarpGroup); + } + cutlass::arch::fence_barrier_init(); + } + + // All producer and consumer threads acquire the initialized barrier state. + __syncthreads(); + + if (is_producer_group) { + cutlass::arch::warpgroup_reg_dealloc(); + if (is_producer_warp) { + const auto& tma_q = parameters_.tma_q; + const auto& tma_k = parameters_.tma_k; + const auto& tma_v = parameters_.tma_v; + const auto& tma_state = parameters_.tma_state; + // TMA and cp.async zero-fill runtime rows [valid_positions, Capacity); + // Q/K/V shared-memory stages contain no fixed [Capacity, MmaM) tail. + auto query_tensor = tma_q.get_tma_tensor( + make_shape(QKVector{}, + QKVectors{}, + parameters.valid_positions, + parameters.hq, + parameters.batch)); + auto key_tensor = tma_k.get_tma_tensor( + make_shape(QKVector{}, + QKVectors{}, + parameters.valid_positions, + parameters.hq, + parameters.batch)); + auto value_tensor = tma_v.get_tma_tensor( + make_shape(BlockDv{}, + parameters.valid_positions, + make_shape(parameters.hv, parameters.batch))); + auto state_tensor = tma_state.get_tma_tensor( + make_shape(WarpThreads{}, + MmaM{}, + StateDvBoxes{}, + StateDkPanels{}, + parameters.heads_per_block * (parameters.state_layer + 1))); + auto state_descriptors = make_tensor( + make_gmem_ptr(parameters.state_tma_descs), + make_layout(make_shape(parameters.num_head_groups, parameters.batch), + make_stride(_1{}, parameters.num_head_groups))); + auto state_head_layout = make_layout( + make_shape(parameters.heads_per_block, parameters.state_layer + 1)); + + auto query_smem = make_tensor( + make_smem_ptr(storage.q.data()), + SmemLayoutQKTma{}); + auto key_smem = make_tensor( + make_smem_ptr(storage.k.data()), + SmemLayoutQKTma{}); + auto value_smem = make_tensor( + make_smem_ptr(storage.value.data()), + SmemLayoutValue{}); + auto state_smem = make_tensor( + make_smem_ptr(reinterpret_cast(storage.state.data())), + SmemLayoutState{}); + + auto query_tma = tma_q.get_slice(_0{}); + auto key_tma = tma_k.get_slice(_0{}); + auto value_tma = tma_v.get_slice(_0{}); + auto state_tma = tma_state.get_slice(_0{}); + + auto tQgQ = query_tma.partition_S(query_tensor); + auto tQsQ = query_tma.partition_D(query_smem); + auto tKgK = key_tma.partition_S(key_tensor); + auto tKsK = key_tma.partition_D(key_smem); + auto tVgV = value_tma.partition_S(value_tensor); + auto tVsV = value_tma.partition_D(value_smem); + auto tSgS = state_tma.partition_S(state_tensor); + auto tSsS = state_tma.partition_D(state_smem); + + auto gate = make_tensor( + make_gmem_ptr(parameters.g), + make_layout(make_shape(parameters.batch, MmaM{}, parameters.hv), + make_stride(parameters.g_batch_stride, + parameters.g_token_stride, + _1{}))); + auto beta = make_tensor( + make_gmem_ptr(parameters.beta), + make_layout(make_shape(parameters.batch, MmaM{}, parameters.hv), + make_stride(parameters.beta_batch_stride, + parameters.beta_token_stride, + _1{}))); + auto gate_smem = make_tensor( + make_smem_ptr(storage.gate.data()), SmemLayoutGate{}); + auto beta_smem = make_tensor( + make_smem_ptr(storage.beta.data()), SmemLayoutBeta{}); + + auto qk_pipe_write = cutlass::make_producer_start_state(); + auto sv_pipe_write = cutlass::make_producer_start_state(); + const int lane_predicate = elect_one_sync(); + auto prefetch_head = [&](const HeadCoordinate& head) { + int commit_length = parameters.valid_positions; + if constexpr (kCommit) { + commit_length = parameters.commit_lengths[head.request]; + if (commit_length == 0 || (parameters.finished && parameters.finished[head.request])) { + return; + } + } + // Every producer lane acquires the reusable Q/K stage before writing it. + qk_pipeline.producer_acquire(qk_pipe_write); + using QKBarrier = typename QKPipeline::ProducerBarrierType; + QKBarrier* qk_barrier = qk_pipeline.producer_get_barrier(qk_pipe_write); + const int write_stage = qk_pipe_write.index(); + if (lane_predicate) { + storage.output_heads[write_stage] = {head.request, head.value_head, commit_length}; + // The producer publishes the decoded output coordinate for this stage. + storage.output_head_ready[write_stage].arrive(); + if constexpr (!kCommit) { + copy(tma_q.with(*qk_barrier), + tQgQ(_, _, _, _0{}, head.query_head, head.request), + tQsQ(_, _, _, _0{}, _0{}, _0{}, write_stage)); + } + copy(tma_k.with(*qk_barrier), + tKgK(_, _, _, _0{}, head.query_head, head.request), + tKsK(_, _, _, _0{}, _0{}, _0{}, write_stage)); + } + auto gate_source = gate(head.request, _, head.value_head); + auto beta_source = beta(head.request, _, head.value_head); + const bool valid = lane < commit_length; + const int source_token = valid ? lane : 0; + const auto gate_lane = GateLaneLayout{}.get_hier_coord(lane); + SM80_CP_ASYNC_CACHEALWAYS_ZFILL::copy( + gate_source(source_token), + gate_smem(get<0>(gate_lane), get<1>(gate_lane), write_stage), + valid); + SM80_CP_ASYNC_CACHEALWAYS_ZFILL::copy( + beta_source(source_token), + beta_smem(get<0>(gate_lane), get<1>(gate_lane), write_stage), + valid); + // Every producer lane commits its gate/beta cp.async completion to the Q/K mbarrier. + qk_pipeline.producer_commit( + qk_pipe_write, cutlass::arch::cpasync_barrier_arrive_noinc); + ++qk_pipe_write; + + // Every producer lane acquires the reusable state/value stage before writing it. + sv_pipeline.producer_acquire(sv_pipe_write); + using SVBarrier = typename SVPipeline::ProducerBarrierType; + SVBarrier* sv_barrier = sv_pipeline.producer_get_barrier(sv_pipe_write); + if (lane_predicate) { + CUTE_UNROLL + for (int half = 0; half < kValueHalves; ++half) { + copy(tma_v.with(*sv_barrier), + tVgV(_, half, _0{}, make_coord(head.value_head, head.request)), + tVsV(_, _0{}, _0{}, half, sv_pipe_write.index())); + } + const TmaDescriptor* state_descriptor = + &state_descriptors(head.head_group, head.request); + CUTE_UNROLL + for (int state_copy = 0; state_copy < size<4>(tSgS); ++state_copy) { + copy(tma_state.with(state_descriptor, *sv_barrier), + tSgS(_, _, _, _, state_copy, + state_head_layout(head.local_head, parameters.state_layer)), + tSsS(_, _, _, _, state_copy, sv_pipe_write.index())); + } + } + ++sv_pipe_write; + }; + + const int persistent_ctas = static_cast(gridDim.x); + const int heads_for_cta = ceil_div( + parameters.hv * parameters.batch - static_cast(blockIdx.x), + persistent_ctas); + auto cta_head_layout = make_layout(make_shape(persistent_ctas, heads_for_cta)); + + CUTLASS_PRAGMA_NO_UNROLL + for (int prefetched_head_in_cta = 0; + prefetched_head_in_cta < heads_for_cta; + ++prefetched_head_in_cta) { + const int prefetched_flattened_head = + int(cta_head_layout(int(blockIdx.x), prefetched_head_in_cta)); + auto prefetched_head = DecodeHead( + prefetched_flattened_head, + parameters.batch, + parameters.hq, + parameters.hv, + parameters.value_heads_per_query_head, + parameters.num_head_groups, + parameters.heads_per_block); + prefetch_head(prefetched_head); + } + // Every producer lane acquires both reusable input stages before publishing termination. + qk_pipeline.producer_acquire(qk_pipe_write); + sv_pipeline.producer_acquire(sv_pipe_write); + if (lane_predicate) { + const int write_stage = qk_pipe_write.index(); + storage.output_heads[write_stage] = {kInvalidRequest, 0, 0}; + // The producer publishes the terminal coordinate after acquiring its reusable stage. + storage.output_head_ready[write_stage].arrive(); + // The elected producer drains both input pipelines after the final head releases them. + qk_pipeline.producer_tail(qk_pipe_write); + sv_pipeline.producer_tail(sv_pipe_write); + } + } + } + else if (warp_group == static_cast(WarpGroupRole::GramPAG)) { + cutlass::arch::warpgroup_reg_alloc(); + auto AG_storage = make_tensor( + make_smem_ptr(storage.AG_bf16.data()), SmemLayoutAG{}); + // K8 _agw_ reads AG[N,K] with WGMMA K=16, while each head writes only + // [0:Capacity,0:Capacity]. Initialize the K tail in every physical stage + // once; the loop is zero-trip for K16. + CUTE_UNROLL + for (int tail_tile = 0; + tail_tile < (kMmaM - Capacity) / kMmaN; + ++tail_tile) { + if (warp_group_thread < size(WgmmaBOperandTailCoordinates{})) { + const auto coordinate = WgmmaBOperandTailCoordinates{}.get_hier_coord( + warp_group_thread); + const int row = int(get<0>(coordinate)); + const int column = Capacity + tail_tile * kMmaN + + int(get<1>(coordinate)); + CUTE_UNROLL + for (int stage = 0; stage < kHandoffStages; ++stage) { + AG_storage(row, column, stage) = Element{}; + } + } + } + QKPipelineState qk_pipe_read; + QKPipelineState qk_pipe_release; + cutlass::PipelineState handoff_stage; + while (true) { + const int read_stage = qk_pipe_read.index(); + // Every compute thread acquires the decoded coordinate before testing the terminal stage. + storage.output_head_ready[read_stage].wait(qk_pipe_read.phase()); + if (storage.output_heads[read_stage].request == kInvalidRequest) { + break; + } + // WG0 acquires Q/K and gate/beta before Gram. + qk_pipeline.consumer_wait(qk_pipe_read); + ProcessGramPAGHead(storage, + read_stage, + handoff_stage.index(), + qk_pipeline, + qk_pipe_release, + lane, + warp, + handoff_stage.phase()); + ++qk_pipe_read; + ++qk_pipe_release; + ++handoff_stage; + } + } + else if (warp_group == static_cast(WarpGroupRole::Output)) { + cutlass::arch::warpgroup_reg_alloc(); + auto P_storage = make_tensor( + make_smem_ptr(storage.P_bf16.data()), SmemLayoutP{}); + auto AGW_storage = make_tensor( + make_smem_ptr(storage.AGW_bf16.data()), SmemLayoutAGW{}); + // K8 _o_final_ reads P[N,K] with WGMMA K=16, while each head writes + // only [0:Capacity,0:Capacity]. Initialize the persistent K tail once; + // the loop is zero-trip for K16. + CUTE_UNROLL + for (int tail_tile = 0; + tail_tile < (kMmaM - Capacity) / kMmaN; + ++tail_tile) { + if (warp_group_thread < size(WgmmaBOperandTailCoordinates{})) { + const auto coordinate = WgmmaBOperandTailCoordinates{}.get_hier_coord( + warp_group_thread); + const int row = int(get<0>(coordinate)); + const int column = Capacity + tail_tile * kMmaN + + int(get<1>(coordinate)); + if constexpr (!kCommit) { + P_storage(row, column) = Element{}; + } + } + } + // K8 _o_final_ reads AGW[Dv,K] with WGMMA K=16, while each head + // writes only columns [0,Capacity). Initialize the local K tail once; + // it persists across heads and this loop is zero-trip for K16. + CUTE_UNROLL + for (int column = Capacity; column < kMmaM; ++column) { + AGW_storage(warp_group_thread, column) = Element{}; + } + QKPipelineState qk_pipe_read; + QKPipelineState qk_pipe_release; + SVPipelineState sv_pipe_read; + SVPipelineState sv_pipe_release; + cutlass::PipelineState handoff_stage; + while (true) { + const int read_stage = qk_pipe_read.index(); + // Every compute thread acquires the decoded coordinate before testing the terminal stage. + storage.output_head_ready[read_stage].wait(qk_pipe_read.phase()); + if (storage.output_heads[read_stage].request == kInvalidRequest) { + break; + } + // WG1 acquires Q/K and gate/beta before P. + qk_pipeline.consumer_wait(qk_pipe_read); + State* commit_state = nullptr; + int commit_length = 0; + if constexpr (kCommit) { + const auto head = storage.output_heads[read_stage]; + commit_length = head.commit_length; + auto state_ptrs = make_tensor( + make_gmem_ptr(parameters.state_ptrs), + make_layout(make_shape(parameters.num_head_groups, parameters.batch), + make_stride(parameters.state_group_stride, parameters.state_request_stride))); + auto state_head_layout = make_layout( + make_shape(parameters.heads_per_block, parameters.num_head_groups)); + const auto state_head = state_head_layout.get_hier_coord(head.value_head); + auto state_heads = make_tensor( + make_gmem_ptr(static_cast(state_ptrs(get<1>(state_head), head.request))), + GmemLayoutState{make_shape(BlockDv{}, HeadDim{}, + parameters.heads_per_block * (parameters.state_layer + 1))}); + auto layer_head_layout = make_layout( + make_shape(parameters.heads_per_block, parameters.state_layer + 1)); + commit_state = &state_heads(_0{}, _0{}, + layer_head_layout(get<0>(state_head), parameters.state_layer)); + } + ProcessFinalHead(storage, + read_stage, + handoff_stage.index(), + qk_pipeline, + qk_pipe_release, + sv_pipeline, + sv_pipe_read, + sv_pipe_release, + parameters.out, + parameters.valid_positions, + parameters.batch, + parameters.hv, + parameters.out_batch_stride, + parameters.out_token_stride, + parameters.out_head_stride, + lane, + warp, + warp_group_thread, + commit_state, + commit_length); + ++qk_pipe_read; + ++qk_pipe_release; + ++sv_pipe_read; + ++sv_pipe_release; + ++handoff_stage; + } + } + else { + cutlass::arch::warpgroup_reg_alloc(); + auto W_storage = make_tensor( + make_smem_ptr(storage.W_bf16.data()), SmemLayoutW{}); + // K8 _agw_ reads W[Dv,K] with WGMMA K=16, while each head writes + // only columns [0,Capacity). Initialize every stage's K tail once; + // it persists across heads and this loop is zero-trip for K16. + CUTE_UNROLL + for (int stage = 0; stage < kHandoffStages; ++stage) { + CUTE_UNROLL + for (int column = Capacity; column < kMmaM; ++column) { + W_storage(warp_group_thread, column, stage) = Element{}; + } + } + QKPipelineState qk_pipe_read; + QKPipelineState qk_pipe_release; + SVPipelineState sv_pipe_read; + SVPipelineState sv_pipe_release; + cutlass::PipelineState handoff_stage; + while (true) { + const int read_stage = qk_pipe_read.index(); + // Every compute thread acquires the decoded coordinate before testing the terminal stage. + storage.output_head_ready[read_stage].wait(qk_pipe_read.phase()); + if (storage.output_heads[read_stage].request == kInvalidRequest) { + break; + } + // WG2 acquires Q/K and gate/beta before prefix. + qk_pipeline.consumer_wait(qk_pipe_read); + ProcessUpdateHead(storage, + read_stage, + handoff_stage.index(), + qk_pipeline, + qk_pipe_release, + sv_pipeline, + sv_pipe_read, + sv_pipe_release, + lane, + warp, + warp_group_thread, + handoff_stage.phase()); + ++qk_pipe_read; + ++qk_pipe_release; + ++sv_pipe_read; + ++sv_pipe_release; + ++handoff_stage; + } + } + } + }; + +public: + static void RegisterStateVariants(Collector& collector) + { + collector.add>(); + collector.add>(); + collector.add>(); + collector.add>(); + } + + const GdrKernelSpec& spec() const noexcept override + { + return kSpec; + } + + const char* name() const noexcept override + { + return kName; + } + + bool Match(const Operation& operation, const PlanningContext& context) const override + { + if (operation.mode == Mode + && (context.token_slots < kMinimumTokens || context.token_slots > spec().chunk_size)) { + return false; + } + Operation selected = operation; + if (selected.mode == Mode && selected.chunk_size == kAutoGdrChunkSize) { + selected.chunk_size = context.token_slots <= kMmaN ? kMmaN : kMmaM; + } + return detail::MatchesGdrSpec(spec(), selected, context); + } + + bool Plan(const Operation& operation, const PlanningContext& context, delta_rule::Plan* plan) const override + { + return detail::PlanSm90Operation(spec(), operation, context, plan); + } + + void PrepareState(const core::Tensor& state_ptrs, + core::Tensor& state_tma_descs, + int layer_groups, + int layers_per_block, + const delta_rule::Plan& plan, + cudaStream_t stream) const override + { + auto* state_prototype_address = reinterpret_cast( + reinterpret_cast(state_ptrs.raw_data()) + & ~uintptr_t(kTmaGlobalAddressAlignment - 1)); + auto state = make_tensor( + make_gmem_ptr(state_prototype_address), + GmemLayoutState{make_shape(BlockDv{}, + HeadDim{}, + plan.problem.heads_per_block * layers_per_block)}); + auto g_state = make_tensor( + state.data(), + select<0, 1, 3, 4, 5>( + flatten(zipped_divide(state.layout(), Tile{})))); + auto tma_state = + make_tma_copy(kTmaLoad, g_state, SmemLayoutStateTmaTile{}, StateTmaTile{}, _1{}); + detail::PrepareSm90StateTmaDescriptors(state_ptrs, + state_tma_descs, + layer_groups, + plan.problem.batch, + plan.problem.num_head_groups, + *tma_state.get_tma_descriptor(), + stream); + } + + void Run(const Arguments& args, const delta_rule::Plan& plan, cudaStream_t stream) const override + { + if constexpr (kCommit) { + if (!args.commit_lengths || args.commit_lengths.dtype() != kInt32 + || args.commit_lengths.device().type != kDEVICE || args.commit_lengths.ndim() != 1 + || args.commit_lengths.shape(0) != plan.problem.batch || args.commit_lengths.stride(0) != 1) { + throw std::invalid_argument("GDR commit_lengths must be a contiguous device int32 [batch] tensor"); + } + if (!args.state_ptrs || (args.state_ptrs.dtype() != kInt64 && args.state_ptrs.dtype() != kPointer) + || args.state_ptrs.device().type != kDEVICE + || (args.state_ptrs.ndim() != 1 && args.state_ptrs.ndim() != 2) + || args.state_ptrs.shape(0) != plan.problem.batch + || (args.state_ptrs.ndim() == 1 && plan.problem.num_head_groups != 1) + || (args.state_ptrs.ndim() == 2 && args.state_ptrs.shape(1) != plan.problem.num_head_groups)) { + throw std::invalid_argument("GDR commit requires device state_ptrs [batch, num_head_groups]"); + } + if (!args.state_tma_descs) { + throw std::invalid_argument("GDR commit requires prepared state_tma_descs"); + } + if (args.finished && (args.finished.dtype() != kBool || args.finished.device().type != kDEVICE + || args.finished.ndim() != 1 || args.finished.shape(0) != plan.problem.batch + || args.finished.stride(0) != 1)) { + throw std::invalid_argument("GDR commit finished must be a contiguous device bool [batch] tensor"); + } + } + using T = Element; + const auto& query = kCommit ? args.k : args.q; + auto g_q = make_tensor( + make_gmem_ptr(reinterpret_cast(query.raw_data())), + make_layout( + make_shape(QKVector{}, + QKVectors{}, + plan.problem.token_num, + plan.problem.hq, + plan.problem.batch), + make_stride(_1{}, + QKVector{}, + query.stride(1), + query.stride(2), + query.stride(0)))); + auto g_k = make_tensor( + make_gmem_ptr(reinterpret_cast(args.k.raw_data())), + make_layout( + make_shape(QKVector{}, + QKVectors{}, + plan.problem.token_num, + plan.problem.hq, + plan.problem.batch), + make_stride(_1{}, + QKVector{}, + args.k.stride(1), + args.k.stride(2), + args.k.stride(0)))); + auto g_v = make_tensor( + make_gmem_ptr(reinterpret_cast(args.v.raw_data())), + make_layout( + make_shape(BlockDv{}, + plan.problem.token_num, + make_shape(plan.problem.hv, plan.problem.batch)), + make_stride(_1{}, + args.v.stride(1), + make_stride(args.v.stride(2), args.v.stride(0))))); + auto tma_k = + make_tma_copy(kTmaLoad, g_k, SmemLayoutQKTmaTile{}, QKTmaTile{}, _1{}); + auto tma_q = tma_k; + if constexpr (!kCommit) { + tma_q = make_tma_copy(kTmaLoad, g_q, SmemLayoutQKTmaTile{}, QKTmaTile{}, _1{}); + } + auto tma_v = + make_tma_copy(kTmaLoad, g_v, SmemLayoutValueTma{}, ValueTmaTile{}, _1{}); + + const int state_layer = + static_cast(args.state_layer_offset + / (int64_t(plan.problem.heads_per_block) * kHeadDim * kBlockDv)); + TmaState tma_state{}; + + const int value_heads_per_query_head = plan.problem.hv / plan.problem.hq; + + using TmaQ = decltype(tma_q); + using TmaK = decltype(tma_k); + using TmaV = decltype(tma_v); + using TmaState = decltype(tma_state); + using Operator = DeviceOperator; + typename Operator::KernelParams parameters{args.g.data(), + args.beta.data(), + reinterpret_cast( + args.state_tma_descs.raw_data()), + kCommit ? nullptr : args.out->data<__nv_bfloat16>(), + plan.problem.token_num, + plan.problem.batch, + plan.problem.hq, + plan.problem.hv, + value_heads_per_query_head, + plan.problem.num_head_groups, + plan.problem.heads_per_block, + args.g.stride(0), + args.g.stride(1), + args.beta.stride(0), + args.beta.stride(1), + kCommit ? 0 : args.out->stride(0), + kCommit ? 0 : args.out->stride(1), + kCommit ? 0 : args.out->stride(2), + state_layer, + kCommit ? reinterpret_cast(args.state_ptrs.raw_data()) : nullptr, + kCommit ? args.commit_lengths.data() : nullptr, + kCommit && args.finished ? args.finished.data() : nullptr, + kCommit ? args.state_ptrs.stride(0) : 0, + kCommit && args.state_ptrs.ndim() == 2 ? args.state_ptrs.stride(1) : 0}; + typename Operator::Params kernel_parameters{tma_q, tma_k, tma_v, tma_state, parameters}; + auto kernel = Sm90GdrVerifyCommitDeviceKernel; + TM_CUDA_CHECK( + cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(kSharedBytes))); + TM_CUDA_CHECK(cudaFuncSetAttribute( + kernel, cudaFuncAttributePreferredSharedMemoryCarveout, cudaSharedmemCarveoutMaxShared)); + + int active_ctas_per_sm = 0; + TM_CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor( + &active_ctas_per_sm, kernel, Policy::kThreads, kSharedBytes)); + TM_CHECK(active_ctas_per_sm > 0) << "SM90 GDR kernel has zero active CTAs"; + + auto head_layout = make_layout(make_shape(plan.problem.hv, plan.problem.batch)); + const int total_heads = int(size(head_layout)); + const int persistent_ctas = std::min(total_heads, getSMCount() * active_ctas_per_sm); + const dim3 grid(persistent_ctas, 1, 1); + const dim3 block(Policy::kThreads); + kernel<<>>(kernel_parameters); + TM_CUDA_CHECK(cudaGetLastError()); + } +}; + +Registrar gdr_reg([](Collector& c) { + Sm90GdrVerifyCommitKernel<8>::RegisterStateVariants(c); + Sm90GdrVerifyCommitKernel<16>::RegisterStateVariants(c); +}); + +} // namespace +} // namespace turbomind::linear_attn::delta_rule diff --git a/src/turbomind/kernels/linear_attn/python_bind.cpp b/src/turbomind/kernels/linear_attn/python_bind.cpp index f86833082d..faf8e9dd6e 100644 --- a/src/turbomind/kernels/linear_attn/python_bind.cpp +++ b/src/turbomind/kernels/linear_attn/python_bind.cpp @@ -10,6 +10,7 @@ #include "src/turbomind/core/tensor.h" #include "src/turbomind/kernels/linear_attn/delta_rule.h" +#include "src/turbomind/kernels/linear_attn/gdn_state_transaction.h" #include "src/turbomind/kernels/linear_attn/registry.h" #include "src/turbomind/utils/cuda_utils.h" @@ -136,6 +137,8 @@ py::dict ProblemToDict(const Problem& problem) out["num_head_groups"] = problem.num_head_groups; out["heads_per_block"] = problem.heads_per_block; out["recurrent"] = IsRecurrentGdr(problem); + out["verify"] = IsVerifyGdr(problem); + out["commit"] = IsCommitGdr(problem); return out; } @@ -146,6 +149,10 @@ const char* ToString(GdrMode mode) return "recurrent"; case GdrMode::kChunked: return "chunked"; + case GdrMode::kVerify: + return "verify"; + case GdrMode::kCommit: + return "commit"; } throw py::value_error("invalid GDR mode"); } @@ -198,7 +205,13 @@ GdrMode ParseMode(const std::string& mode) if (mode == "chunked") { return GdrMode::kChunked; } - throw py::value_error("mode must be one of: recurrent, chunked"); + if (mode == "verify") { + return GdrMode::kVerify; + } + if (mode == "commit") { + return GdrMode::kCommit; + } + throw py::value_error("mode must be one of: recurrent, chunked, verify, commit"); } DataType ParseStateDtype(const std::string& dtype) @@ -265,11 +278,13 @@ Arguments MakeExecutionArguments(const py::object& q, const py::object& state_tma_descs, const py::object& q_offsets, const py::object& finished, + const py::object& commit_lengths, + GdrMode mode, int64_t state_layer_offset, std::vector& storage) { Arguments args{}; - args.q = TensorFromObject(q, "q", true); + args.q = TensorFromObject(q, "q", mode != GdrMode::kCommit); args.k = TensorFromObject(k, "k", true); args.v = TensorFromObject(v, "v", true); args.g = TensorFromObject(g, "g", true); @@ -278,6 +293,7 @@ Arguments MakeExecutionArguments(const py::object& q, args.state_tma_descs = TensorFromObject(state_tma_descs, "state_tma_descs", false); args.q_offsets = TensorFromObject(q_offsets, "q_offsets", false); args.finished = TensorFromObject(finished, "finished", false); + args.commit_lengths = TensorFromObject(commit_lengths, "commit_lengths", mode == GdrMode::kCommit); args.out = OptionalTensorPtr(out, "out", storage); args.workspace = OptionalTensorPtr(workspace, "workspace", storage); args.state_layer_offset = state_layer_offset; @@ -320,7 +336,8 @@ py::dict PlanBridge(const py::object& q, int num_head_groups, int heads_per_block) { - const auto q_tensor = TensorFromObject(q, "q", true); + const auto parsed_mode = ParseMode(mode); + const auto q_tensor = TensorFromObject(q, "q", parsed_mode != GdrMode::kCommit); const auto k_tensor = TensorFromObject(k, "k", true); const auto v_tensor = TensorFromObject(v, "v", true); const auto g_tensor = TensorFromObject(g, "g", true); @@ -329,10 +346,9 @@ py::dict PlanBridge(const py::object& q, static_cast(k_tensor); static_cast(beta_tensor); - const auto parsed_mode = ParseMode(mode); const auto parsed_state_dtype = ParseStateDtype(state_dtype); const auto operation = MakeOperation(parsed_mode, chunk_size, ParseContextParallelLevel(cp_level)); - const auto context = MakePlanningContext(q_tensor, + const auto context = MakePlanningContext(parsed_mode == GdrMode::kCommit ? k_tensor : q_tensor, v_tensor, g_tensor, beta_tensor, @@ -400,7 +416,8 @@ void RunBridge(const py::object& q, const py::object& state_tma_descs, const py::object& q_offsets, const py::object& finished, - int64_t state_layer_offset) + int64_t state_layer_offset, + const py::object& commit_lengths) { auto* plan_ptr = PlanFromDict(plan); std::vector storage; @@ -416,6 +433,8 @@ void RunBridge(const py::object& q, state_tma_descs, q_offsets, finished, + commit_lengths, + plan_ptr->problem.mode, state_layer_offset, storage); GatedDeltaRule rule; @@ -427,6 +446,115 @@ void RunBridge(const py::object& q, } } +void BuildStateStoreMaskBridge(const py::object& out, + const py::object& finished, + const py::object& speculative, + std::uintptr_t stream_ptr) +{ + auto out_tensor = TensorFromObject(out, "out", true); + auto finished_tensor = TensorFromObject(finished, "finished", true); + auto speculative_tensor = TensorFromObject(speculative, "speculative", true); + invokeBuildGdnStateStoreMask(out_tensor.data(), + finished_tensor.data(), + speculative_tensor.data(), + static_cast(out_tensor.size()), + reinterpret_cast(stream_ptr)); +} + +void CaptureTransitionsBridge(const py::object& raw_projection, + const py::object& normalized_key, + const py::object& value, + const py::object& log_decay, + const py::object& beta, + const py::object& q_offsets, + const py::object& request_indices, + int gdn_layer, + int verify_positions, + const py::object& journal_raw, + const py::object& journal_key, + const py::object& journal_value, + const py::object& journal_decay, + const py::object& journal_beta, + std::uintptr_t stream_ptr) +{ + TransitionJournal journal{TensorFromObject(journal_raw, "journal_raw", true), + TensorFromObject(journal_key, "journal_key", true), + TensorFromObject(journal_value, "journal_value", true), + TensorFromObject(journal_decay, "journal_decay", true), + TensorFromObject(journal_beta, "journal_beta", true)}; + invokeCaptureGdnTransitions(TensorFromObject(raw_projection, "raw_projection", true), + TensorFromObject(normalized_key, "normalized_key", true), + TensorFromObject(value, "value", true), + TensorFromObject(log_decay, "log_decay", true), + TensorFromObject(beta, "beta", true), + TensorFromObject(q_offsets, "q_offsets", true).buffer(), + TensorFromObject(request_indices, "request_indices", true).buffer(), + gdn_layer, + verify_positions, + std::move(journal), + reinterpret_cast(stream_ptr)); +} + +void CommitConvStateBridge(const py::object& raw_conv, + const py::object& conv_state_ptrs, + const py::object& request_indices, + const py::object& entry_sequence_length, + const py::object& accept_len, + const py::object& conv_state_offsets, + int conv_dim, + int d_conv, + std::uintptr_t stream_ptr) +{ + auto pointers = TensorFromObject(conv_state_ptrs, "conv_state_ptrs", true); + Buffer_ pointer_buffer{static_cast(pointers.raw_data()), pointers.size(), pointers.device()}; + invokeCommitAcceptedConvState(TensorFromObject(raw_conv, "raw_conv", true), + pointer_buffer, + TensorFromObject(request_indices, "request_indices", true).buffer(), + TensorFromObject(entry_sequence_length, "entry_sequence_length", true).buffer(), + TensorFromObject(accept_len, "accept_len", true).buffer(), + TensorFromObject(conv_state_offsets, "conv_state_offsets", true).buffer(), + conv_dim, + d_conv, + reinterpret_cast(stream_ptr)); +} + +void CommitRecurrentStateBridge(const py::object& key, + const py::object& value, + const py::object& log_decay, + const py::object& beta, + const py::object& recurrent_state_ptrs, + const py::object& request_indices, + const py::object& accept_len, + const std::string& state_dtype, + int layers_per_block, + int heads_per_block, + std::uintptr_t stream_ptr) +{ + auto key_tensor = TensorFromObject(key, "key", true); + auto value_tensor = TensorFromObject(value, "value", true); + auto pointers = TensorFromObject(recurrent_state_ptrs, "recurrent_state_ptrs", true); + Tensor pointer_tensor{pointers.raw_data(), pointers.layout(), data_type_v, pointers.device()}; + + AcceptedPrefixArguments args{}; + args.key = key_tensor; + args.value = value_tensor; + args.log_decay = TensorFromObject(log_decay, "log_decay", true); + args.beta = TensorFromObject(beta, "beta", true); + args.recurrent_state_ptrs = std::move(pointer_tensor); + args.request_indices = TensorFromObject(request_indices, "request_indices", true); + args.accept_len = TensorFromObject(accept_len, "accept_len", true); + args.layer_count = static_cast(key_tensor.shape(0)); + args.speculative_count = static_cast(key_tensor.shape(1)); + args.position_count = static_cast(key_tensor.shape(2)); + args.hq = static_cast(key_tensor.shape(3)); + args.hv = static_cast(value_tensor.shape(3)); + args.num_head_groups = static_cast(pointers.shape(2)); + args.layers_per_block = layers_per_block; + args.heads_per_block = heads_per_block; + args.sm_count = getSMCount(); + invokeCommitAcceptedRecurrentState(args, ParseStateDtype(state_dtype), reinterpret_cast(stream_ptr)); +} + } // namespace void bind_delta_rule(py::module_& module) @@ -470,7 +598,56 @@ void bind_delta_rule(py::module_& module) "state_tma_descs"_a = py::none(), "q_offsets"_a = py::none(), "finished"_a = py::none(), - "state_layer_offset"_a = int64_t{0}); + "state_layer_offset"_a = int64_t{0}, + "commit_lengths"_a = py::none()); + + module.def("gdn_build_state_store_mask", + &BuildStateStoreMaskBridge, + "out"_a, + "finished"_a, + "speculative"_a, + "stream_ptr"_a = std::uintptr_t{0}); + module.def("gdn_capture_transitions", + &CaptureTransitionsBridge, + "raw_projection"_a, + "normalized_key"_a, + "value"_a, + "log_decay"_a, + "beta"_a, + "q_offsets"_a, + "request_indices"_a, + "gdn_layer"_a, + "verify_positions"_a, + "journal_raw"_a, + "journal_key"_a, + "journal_value"_a, + "journal_decay"_a, + "journal_beta"_a, + "stream_ptr"_a = std::uintptr_t{0}); + module.def("gdn_commit_conv_state", + &CommitConvStateBridge, + "raw_conv"_a, + "conv_state_ptrs"_a, + "request_indices"_a, + "entry_sequence_length"_a, + "accept_len"_a, + "conv_state_offsets"_a, + "conv_dim"_a, + "d_conv"_a, + "stream_ptr"_a = std::uintptr_t{0}); + module.def("gdn_commit_recurrent_state", + &CommitRecurrentStateBridge, + "key"_a, + "value"_a, + "log_decay"_a, + "beta"_a, + "recurrent_state_ptrs"_a, + "request_indices"_a, + "accept_len"_a, + "state_dtype"_a, + "layers_per_block"_a, + "heads_per_block"_a, + "stream_ptr"_a = std::uintptr_t{0}); } } // namespace turbomind::linear_attn::delta_rule diff --git a/src/turbomind/kernels/norm/rms_norm.cu b/src/turbomind/kernels/norm/rms_norm.cu index d47e317a40..f2682f19e9 100644 --- a/src/turbomind/kernels/norm/rms_norm.cu +++ b/src/turbomind/kernels/norm/rms_norm.cu @@ -135,6 +135,30 @@ void invokeRMSNorm(Tensor& out, const Tensor& x, const Tensor& w, float eps, boo TM_CUDA_CHECK(cudaGetLastError()); } +void invokeRMSNormConcat(Tensor& out, + const Tensor& left, + const Tensor& left_weight, + float left_eps, + bool left_zero_centered, + const Tensor& right, + const Tensor& right_weight, + float right_eps, + bool right_zero_centered, + cudaStream_t stream) +{ + const auto [token_num, hidden_size] = left.shapes(0, 1); + + if (token_num == 0) { + return; + } + + Tensor left_out = out.slice({0, 0}, {token_num, hidden_size}); + Tensor right_out = out.slice({0, hidden_size}, {token_num, hidden_size}); + + invokeRMSNorm(left_out, left, left_weight, left_eps, left_zero_centered, stream); + invokeRMSNorm(right_out, right, right_weight, right_eps, right_zero_centered, stream); +} + namespace kernel { template diff --git a/src/turbomind/kernels/norm/rms_norm.h b/src/turbomind/kernels/norm/rms_norm.h index b50074994c..1ff3442833 100644 --- a/src/turbomind/kernels/norm/rms_norm.h +++ b/src/turbomind/kernels/norm/rms_norm.h @@ -17,6 +17,17 @@ void invokeQkRMSNorm(Tensor& qkv, bool zero_centered, cudaStream_t st); +void invokeRMSNormConcat(Tensor& out, + const Tensor& left, + const Tensor& left_weight, + float left_eps, + bool left_zero_centered, + const Tensor& right, + const Tensor& right_weight, + float right_eps, + bool right_zero_centered, + cudaStream_t stream); + template void invokeBiasResidualRMSNorm(T* residual, T* hidden_states, diff --git a/src/turbomind/kernels/sampling_device.cuh b/src/turbomind/kernels/sampling_device.cuh new file mode 100644 index 0000000000..a9522b2b6b --- /dev/null +++ b/src/turbomind/kernels/sampling_device.cuh @@ -0,0 +1,55 @@ +#pragma once + +#ifndef CUDART_VERSION +#error CUDART_VERSION Undefined! +#elif (CUDART_VERSION >= 11000) +#include +#else +#include "3rdparty/cub/cub.cuh" +#endif + +#include + +#include "src/turbomind/kernels/sampling_topp_kernels.h" + +namespace turbomind { + +template +struct ProcessedDistributionSampleStorage { + typename cub::BlockScan::TempStorage scan; + float threshold; + int selected_index; +}; + +template +__device__ int SampleProcessedDistribution(const T* row, + int kept, + curandState_t* random_state, + ProcessedDistributionSampleStorage& storage) +{ + const int tid = threadIdx.x; + if (tid == 0) { + storage.threshold = curand_uniform(random_state); + } + __syncthreads(); + + BlockPrefixCallbackOp prefix_op{0.f}; + const int end = (kept + BlockSize - 1) / BlockSize * BlockSize; + for (int i = tid; i < end; i += BlockSize) { + const float probability = i < kept ? static_cast(row[i]) : 0.f; + float inclusive_mass{}; + cub::BlockScan(storage.scan).InclusiveSum(probability, inclusive_mass, prefix_op); + + const int count = __syncthreads_count(inclusive_mass > storage.threshold); + if (count != 0 || i + BlockSize >= end) { + if (tid == min(BlockSize - count, BlockSize - 1)) { + storage.selected_index = min(i, kept - 1); + } + break; + } + } + __syncthreads(); + return storage.selected_index; +} + +} // namespace turbomind diff --git a/src/turbomind/kernels/sampling_kernels.cu b/src/turbomind/kernels/sampling_kernels.cu index b42d2d0e8a..a762ae48b5 100644 --- a/src/turbomind/kernels/sampling_kernels.cu +++ b/src/turbomind/kernels/sampling_kernels.cu @@ -1,67 +1,40 @@ -#ifndef CUDART_VERSION -#error CUDART_VERSION Undefined! -#elif (CUDART_VERSION >= 11000) -#include -#else -#include "3rdparty/cub/cub.cuh" -#endif +#include "src/turbomind/kernels/sampling_device.cuh" #include "src/turbomind/kernels/sampling_kernels.h" -#include "src/turbomind/kernels/sampling_topp_kernels.h" #include "src/turbomind/utils/constant.h" #include "src/turbomind/utils/cuda_utils.h" namespace turbomind { template -__global__ void sampling(const T* logits, +__global__ void sampling(const T* probabilities, const int stride, const int* indices, const int* kept, curandState_t* curandstate, const int* curandstate_indices, - int* output_ids, - int* sequence_length, + const bool* sample_mask, + int* selected_tokens, T* sampled_logprobs, int* sampled_indexes, int* sampled_nums) { - int tid = threadIdx.x; - int batch_id = blockIdx.x; - int n = kept[batch_id]; + const int batch_id = blockIdx.x; - logits += stride * batch_id; - indices += stride * batch_id; - - __shared__ float rand_num_s; - __shared__ int selected; - if (tid == 0) { - const int state_row = curandstate_indices[batch_id]; - rand_num_s = curand_uniform(curandstate + state_row); + if (sample_mask != nullptr && !sample_mask[batch_id]) { + return; } - __syncthreads(); - typedef cub::BlockScan BlockScan; - __shared__ typename BlockScan::TempStorage temp_storage; + const int tid = threadIdx.x; + const int n = kept[batch_id]; - float local_rand = rand_num_s; - float prefix_sum = 0.f; - BlockPrefixCallbackOp prefix_op{0}; - int end = (n + BLOCK_SIZE - 1) / BLOCK_SIZE * BLOCK_SIZE; - for (int i = tid; i < end; i += BLOCK_SIZE) { - float thread_logit = (i < n) ? static_cast(logits[i]) : 0.f; - BlockScan(temp_storage).InclusiveSum(thread_logit, prefix_sum, prefix_op); - auto count = __syncthreads_count(prefix_sum > local_rand); - if (count != 0 || (i + BLOCK_SIZE) >= end) { - if (tid == min(BLOCK_SIZE - count, BLOCK_SIZE - 1)) { - selected = min(i, n - 1); - output_ids[batch_id] = indices[selected]; - } - break; - } - } + probabilities += stride * batch_id; + indices += stride * batch_id; + __shared__ ProcessedDistributionSampleStorage storage; + const int selected = SampleProcessedDistribution( + probabilities, n, curandstate + curandstate_indices[batch_id], storage); if (tid == 0) { - sequence_length[batch_id] += 1; + selected_tokens[batch_id] = indices[selected]; } if (sampled_logprobs != nullptr && sampled_indexes != nullptr && sampled_nums != nullptr) { @@ -70,12 +43,12 @@ __global__ void sampling(const T* logits, sampled_indexes += batch_id * kMaxLogProb; int end = min(n, kMaxLogProb); for (int i = tid; i < end; i += BLOCK_SIZE) { - sampled_logprobs[i] = logf(logits[i]); + sampled_logprobs[i] = logf(probabilities[i]); sampled_indexes[i] = indices[i]; } if (n > kMaxLogProb && selected >= kMaxLogProb) { if ((kMaxLogProb - 1 + BLOCK_SIZE - tid) % BLOCK_SIZE == 0) { - sampled_logprobs[kMaxLogProb - 1] = logf(logits[selected]); + sampled_logprobs[kMaxLogProb - 1] = logf(probabilities[selected]); sampled_indexes[kMaxLogProb - 1] = indices[selected]; } } @@ -86,16 +59,20 @@ __global__ void sampling(const T* logits, template void invokeSampling(SamplingParams& params, cudaStream_t stream) { + if (params.batch_size == 0) { + return; + } + const int grid = params.batch_size; const int block = 256; - sampling<<>>((T*)params.logits, + sampling<<>>((const T*)params.probabilities, params.stride, params.indices, params.kept, params.curandstate, params.curandstate_indices, - params.output_ids, - params.sequence_length, + params.sample_mask, + params.selected_tokens, (T*)params.sampled_logprobs, params.sampled_indexes, params.sampled_nums); diff --git a/src/turbomind/kernels/sampling_kernels.h b/src/turbomind/kernels/sampling_kernels.h index c3822fc329..ab8c250e6c 100644 --- a/src/turbomind/kernels/sampling_kernels.h +++ b/src/turbomind/kernels/sampling_kernels.h @@ -24,15 +24,15 @@ namespace turbomind { struct SamplingParams { - void* logits; + const void* probabilities; int stride; - int* indices; - int* kept; + const int* indices; + const int* kept; curandState_t* curandstate; const int* curandstate_indices; + const bool* sample_mask; size_t batch_size; - int* output_ids; - int* sequence_length; + int* selected_tokens; void* sampled_logprobs; int* sampled_indexes; int* sampled_nums; diff --git a/src/turbomind/kernels/sampling_penalty_kernels.cu b/src/turbomind/kernels/sampling_penalty_kernels.cu index e63faf07c9..0ac44c179a 100644 --- a/src/turbomind/kernels/sampling_penalty_kernels.cu +++ b/src/turbomind/kernels/sampling_penalty_kernels.cu @@ -137,11 +137,16 @@ __global__ void RepetitionPenaltyKernel(T* logits, const float* penalties, const int* const* token_ids_ptrs, const int* sequence_length, + const bool* logits_active, int vocab_size, int mask_size) { const int bi = blockIdx.x; + if (logits_active != nullptr && !logits_active[bi]) { + return; + } + const int seq_len = sequence_length[bi]; const int* token_ids = token_ids_ptrs[bi]; @@ -176,6 +181,7 @@ void ApplyRepetitionPenalty(Tensor& logits, const Buffer_& penalties, const Buffer_& token_ids_ptrs, const Buffer_& sequence_length, + const bool* logits_active, cudaStream_t stream) { TM_CHECK_EQ(logits.ndim(), 2); @@ -189,8 +195,13 @@ void ApplyRepetitionPenalty(Tensor& logits, TM_CHECK_EQ(cudaFuncSetAttribute(func, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size), 0); } TM_LOG_DEBUG("smem_size = {}", smem_size); - func<<>>( - logits.data(), penalties.data(), token_ids_ptrs.data(), sequence_length.data(), vocab_size, mask_size); + func<<>>(logits.data(), + penalties.data(), + token_ids_ptrs.data(), + sequence_length.data(), + logits_active, + vocab_size, + mask_size); }; invoke(float{}); TM_CUDA_CHECK(cudaGetLastError()); diff --git a/src/turbomind/kernels/sampling_penalty_kernels.h b/src/turbomind/kernels/sampling_penalty_kernels.h index 05992eb4fe..72aad6ae83 100644 --- a/src/turbomind/kernels/sampling_penalty_kernels.h +++ b/src/turbomind/kernels/sampling_penalty_kernels.h @@ -27,6 +27,7 @@ void ApplyRepetitionPenalty(Tensor& logits, const Buffer_& penalties, const Buffer_& token_ids_ptrs, const Buffer_& sequence_length, + const bool* logits_active, cudaStream_t stream); template diff --git a/src/turbomind/kernels/sampling_topp_kernels.cu b/src/turbomind/kernels/sampling_topp_kernels.cu index ace05d727b..d9d2e4546e 100644 --- a/src/turbomind/kernels/sampling_topp_kernels.cu +++ b/src/turbomind/kernels/sampling_topp_kernels.cu @@ -33,6 +33,41 @@ namespace turbomind { +namespace { + +template +size_t GetTopPSortCubBytes(int batch_size, int vocab_size, int vocab_size_padded, cudaStream_t stream) +{ + const int num_items = vocab_size_padded * (batch_size - 1) + vocab_size; + + size_t cub_bytes{}; + cub::DeviceSegmentedRadixSort::SortPairsDescending(nullptr, + cub_bytes, + static_cast(nullptr), + static_cast(nullptr), + static_cast(nullptr), + static_cast(nullptr), + num_items, + batch_size, + static_cast(nullptr), + static_cast(nullptr), + 0, + sizeof(T) * 8, + stream); + + return cub_bytes; +} + +} // namespace + +size_t GetTopPSortWorkspaceBytes(int batch_size, int vocab_size, int vocab_size_padded, cudaStream_t stream) +{ + const size_t item_count = static_cast(batch_size) * vocab_size_padded; + + return GetTopPSortCubBytes(batch_size, vocab_size, vocab_size_padded, stream) + item_count * sizeof(int) + + 2 * static_cast(batch_size) * sizeof(int); +} + __global__ void topPSortInitialize(const int vocab_size_padded, const int vocab_size, const size_t batch_size, @@ -222,20 +257,8 @@ void invokeTopPSort(TopPSortParams& params, cudaStream_t stream) { const int num_items = params.vocab_size_padded * (params.batch_size - 1) + params.vocab_size; - size_t cub_temp_storage_size{}; - TM_CUDA_CHECK(cub::DeviceSegmentedRadixSort::SortPairsDescending(nullptr, - cub_temp_storage_size, - (T*)nullptr, - (T*)nullptr, - (int*)nullptr, - (int*)nullptr, - num_items, - params.batch_size, - (int*)nullptr, - (int*)nullptr, - 0, // begin_bit - sizeof(T) * 8, // end_bit = sizeof(KeyT) * 8 - stream)); // cudaStream_t + size_t cub_temp_storage_size = + GetTopPSortCubBytes(params.batch_size, params.vocab_size, params.vocab_size_padded, stream); TM_CHECK(core::Context::stream().handle() == stream); diff --git a/src/turbomind/kernels/sampling_topp_kernels.h b/src/turbomind/kernels/sampling_topp_kernels.h index ca868b3ffe..f68b1f1913 100644 --- a/src/turbomind/kernels/sampling_topp_kernels.h +++ b/src/turbomind/kernels/sampling_topp_kernels.h @@ -15,10 +15,14 @@ */ #pragma once +#include + #include namespace turbomind { +size_t GetTopPSortWorkspaceBytes(int batch_size, int vocab_size, int vocab_size_padded, cudaStream_t stream); + void invokeTopPSortInitialize(const int vocab_size_padded, const int vocab_size, const size_t batch_size, diff --git a/src/turbomind/kernels/speculative_sampling_kernels.cu b/src/turbomind/kernels/speculative_sampling_kernels.cu new file mode 100644 index 0000000000..b41fd2e8fd --- /dev/null +++ b/src/turbomind/kernels/speculative_sampling_kernels.cu @@ -0,0 +1,308 @@ +#include +#include +#include + +#ifndef CUDART_VERSION +#error CUDART_VERSION Undefined! +#elif (CUDART_VERSION >= 11000) +#include +#else +#include "3rdparty/cub/cub.cuh" +#endif + +#include + +#include "src/turbomind/kernels/sampling_device.cuh" +#include "src/turbomind/kernels/speculative_sampling_kernels.h" + +namespace turbomind { +namespace { + +constexpr int kBlockSize = 256; + +struct ArgMax { + float probability; + int index; +}; + +struct ArgMaxOp { + __device__ ArgMax operator()(ArgMax lhs, ArgMax rhs) const + { + if (rhs.probability > lhs.probability || (rhs.probability == lhs.probability && rhs.index < lhs.index)) { + return rhs; + } + return lhs; + } +}; + +struct RecoveryStats { + float mass; + int last_index; +}; + +struct RecoveryStatsOp { + __device__ RecoveryStats operator()(RecoveryStats lhs, RecoveryStats rhs) const + { + return {lhs.mass + rhs.mass, max(lhs.last_index, rhs.last_index)}; + } +}; + +struct MaxFloatOp { + __device__ float operator()(float lhs, float rhs) const + { + return max(lhs, rhs); + } +}; + +struct MinIndexOp { + __device__ int operator()(int lhs, int rhs) const + { + return min(lhs, rhs); + } +}; + +template +struct VerifySharedStorage { + typename cub::BlockReduce::TempStorage argmax; + typename cub::BlockReduce::TempStorage float_reduce; + typename cub::BlockReduce::TempStorage recovery; + typename cub::BlockScan::TempStorage recovery_scan; + typename cub::BlockReduce::TempStorage index_reduce; + ProcessedDistributionSampleStorage sample; + + float probability_draft; + float recovery_threshold; + int recovery_token; + bool recovery_found; +}; + +template +__device__ void VerifyOneProcessedDraft(const float* row, + const int* token_row, + int kept, + int draft, + bool greedy, + curandState_t* random_state, + int& selected, + bool& accepted, + VerifySharedStorage& storage) +{ + const int thread_idx = threadIdx.x; + + if (greedy) { + ArgMax local{-FLT_MAX, INT_MAX}; + for (int i = thread_idx; i < kept; i += BlockSize) { + local = ArgMaxOp{}(local, ArgMax{row[i], i}); + } + + const ArgMax target = cub::BlockReduce(storage.argmax).Reduce(local, ArgMaxOp{}); + if (thread_idx == 0) { + selected = token_row[target.index]; + accepted = selected == draft; + } + __syncthreads(); + return; + } + + float local_draft_probability = 0.f; + for (int i = thread_idx; i < kept; i += BlockSize) { + if (token_row[i] == draft) { + local_draft_probability = max(local_draft_probability, row[i]); + } + } + const float reduced_draft_probability = + cub::BlockReduce(storage.float_reduce).Reduce(local_draft_probability, MaxFloatOp{}); + if (thread_idx == 0) { + storage.probability_draft = reduced_draft_probability; + } + __syncthreads(); + + if (thread_idx == 0) { + const float acceptance_uniform = curand_uniform(random_state); + accepted = storage.probability_draft == 1.f + || (storage.probability_draft > 0.f && acceptance_uniform <= storage.probability_draft); + if (accepted) { + selected = draft; + } + } + __syncthreads(); + if (accepted) { + return; + } + + RecoveryStats local_recovery{0.f, -1}; + for (int i = thread_idx; i < kept; i += BlockSize) { + const float probability = row[i]; + if (token_row[i] != draft && probability > 0.f) { + local_recovery.mass += probability; + local_recovery.last_index = max(local_recovery.last_index, i); + } + } + const RecoveryStats recovery = + cub::BlockReduce(storage.recovery).Reduce(local_recovery, RecoveryStatsOp{}); + if (thread_idx == 0) { + storage.recovery_threshold = curand_uniform(random_state) * recovery.mass; + storage.recovery_token = token_row[recovery.last_index]; + storage.recovery_found = false; + } + __syncthreads(); + + BlockPrefixCallbackOp prefix_op{0.f}; + for (int base = 0; base < kept; base += BlockSize) { + const int i = base + thread_idx; + float probability = 0.f; + if (i < kept && token_row[i] != draft && row[i] > 0.f) { + probability = row[i]; + } + + float inclusive_mass; + cub::BlockScan(storage.recovery_scan).InclusiveSum(probability, inclusive_mass, prefix_op); + __syncthreads(); + + const int candidate = + probability > 0.f && inclusive_mass > storage.recovery_threshold ? i : INT_MAX; + const int first_crossing = + cub::BlockReduce(storage.index_reduce).Reduce(candidate, MinIndexOp{}); + if (thread_idx == 0 && first_crossing != INT_MAX) { + storage.recovery_token = token_row[first_crossing]; + storage.recovery_found = true; + } + __syncthreads(); + if (storage.recovery_found) { + break; + } + } + + if (thread_idx == 0) { + selected = storage.recovery_token; + accepted = false; + } + __syncthreads(); +} + +template +__device__ void SampleOneProcessedDistribution(const float* row, + const int* token_row, + int kept, + curandState_t* random_state, + int& selected, + VerifySharedStorage& storage) +{ + const int selected_index = + SampleProcessedDistribution(row, kept, random_state, storage.sample); + if (threadIdx.x == 0) { + selected = token_row[selected_index]; + } + __syncthreads(); +} + +template +__global__ void VerifyTargetBlock(VerifyTargetBlockParams p) +{ + const int b = blockIdx.x; + const int g_begin = p.request_to_generation_offsets[b]; + const int g_end = p.request_to_generation_offsets[b + 1]; + if (g_begin == g_end) { + return; + } + + const int g = g_begin; + __shared__ VerifySharedStorage storage; + __shared__ int selected; + __shared__ int span_len; + __shared__ int accepted_drafts; + __shared__ bool accepted; + __shared__ bool continue_verification; + + curandState_t* random_state = p.random_states + p.random_state_indices[g]; + if (threadIdx.x == 0) { + span_len = 0; + accepted_drafts = 0; + continue_verification = p.logits_active[g]; + } + __syncthreads(); + + if (!p.speculative_row[b]) { + if (continue_verification) { + const float* row = p.probabilities + g * p.probability_stride; + const int* ids = p.probability_token_ids + g * p.token_id_stride; + SampleOneProcessedDistribution( + row, ids, p.kept_count[g], random_state, selected, storage); + __syncthreads(); + + if (threadIdx.x == 0) { + p.request_token_ids_ptrs[b][p.entry_sequence_length[b]] = selected; + p.selected_span_ids[b * p.selected_span_stride] = selected; + span_len = 1; + } + } + } + else { + for (int position = 0; position < p.position_count - 1; ++position) { + const int flat = position * p.generation_count + g; + if (continue_verification && p.logits_active[flat]) { + const float* row = p.probabilities + flat * p.probability_stride; + const int* ids = p.probability_token_ids + flat * p.token_id_stride; + const int draft = p.verification_draft_ids[position * p.draft_row_stride + g]; + + VerifyOneProcessedDraft(row, + ids, + p.kept_count[flat], + draft, + p.greedy[flat], + random_state, + selected, + accepted, + storage); + __syncthreads(); + + if (threadIdx.x == 0) { + p.request_token_ids_ptrs[b][p.entry_sequence_length[b] + span_len] = selected; + p.selected_span_ids[b * p.selected_span_stride + span_len] = selected; + ++span_len; + if (accepted) { + ++accepted_drafts; + } + else { + continue_verification = false; + } + } + __syncthreads(); + } + } + + const int bonus_flat = (p.position_count - 1) * p.generation_count + g; + if (continue_verification && p.logits_active[bonus_flat]) { + const float* row = p.probabilities + bonus_flat * p.probability_stride; + const int* ids = p.probability_token_ids + bonus_flat * p.token_id_stride; + SampleOneProcessedDistribution( + row, ids, p.kept_count[bonus_flat], random_state, selected, storage); + __syncthreads(); + + if (threadIdx.x == 0) { + p.request_token_ids_ptrs[b][p.entry_sequence_length[b] + span_len] = selected; + p.selected_span_ids[b * p.selected_span_stride + span_len] = selected; + ++span_len; + } + } + } + + if (threadIdx.x == 0) { + p.accept_len[b] = span_len; + if (p.accepted_draft_count && p.speculative_row[b] && p.logits_active[g]) { + p.accepted_draft_count[b] = accepted_drafts; + } + } +} + +} // namespace + +void invokeVerifyTargetBlock(const VerifyTargetBlockParams& params, cudaStream_t stream) +{ + if (params.request_count == 0) { + return; + } + VerifyTargetBlock<<>>(params); +} + +} // namespace turbomind diff --git a/src/turbomind/kernels/speculative_sampling_kernels.h b/src/turbomind/kernels/speculative_sampling_kernels.h new file mode 100644 index 0000000000..666aec7b4f --- /dev/null +++ b/src/turbomind/kernels/speculative_sampling_kernels.h @@ -0,0 +1,39 @@ +#pragma once + +#include +#include + +namespace turbomind { + +struct VerifyTargetBlockParams { + const float* probabilities; + int probability_stride; + const int* probability_token_ids; + int token_id_stride; + const int* kept_count; + const int* verification_draft_ids; + int draft_row_stride; + const bool* greedy; + const bool* logits_active; + + curandState_t* random_states; + const int* random_state_indices; + + int* const* request_token_ids_ptrs; + const int* entry_sequence_length; + const int* request_to_generation_offsets; + const bool* speculative_row; + + int* selected_span_ids; + int selected_span_stride; + int* accept_len; + int* accepted_draft_count; + + int request_count; + int generation_count; + int position_count; +}; + +void invokeVerifyTargetBlock(const VerifyTargetBlockParams& params, cudaStream_t stream); + +} // namespace turbomind diff --git a/src/turbomind/kernels/speculative_sampling_python_bind.cpp b/src/turbomind/kernels/speculative_sampling_python_bind.cpp new file mode 100644 index 0000000000..8f272d0ff8 --- /dev/null +++ b/src/turbomind/kernels/speculative_sampling_python_bind.cpp @@ -0,0 +1,240 @@ +#include +#include + +#include +#include + +#include + +#include "src/turbomind/kernels/sampling_kernels.h" +#include "src/turbomind/kernels/sampling_topk_kernels.h" +#include "src/turbomind/kernels/speculative_sampling_kernels.h" +#include "src/turbomind/models/llama/llama_kernels.h" +#include "src/turbomind/python/eagle3_component_bindings.h" +#include "src/turbomind/python/eagle3_dlpack_internal.h" +#include "src/turbomind/utils/cuda_utils.h" + +namespace py = pybind11; + +namespace turbomind::python { +namespace { + +int GetCudaOrdinal(py::handle tensor) +{ + return tensor.attr("__dlpack_device__")().cast()[1].cast(); +} + +} // namespace + +void BindSpeculativeSampling(py::module_& module) +{ + module.def( + "initialize_speculative_sampling_states", + [](py::handle random_states_object, + int random_state_count, + py::handle random_seeds_object, + py::handle initialize_object, + uintptr_t stream_ptr) { + CudaDeviceGuard guard{GetCudaOrdinal(random_states_object)}; + auto random_states = detail::ConsumeDLPackWithStrides(random_states_object, stream_ptr); + auto random_seeds = detail::ConsumeDLPackWithStrides(random_seeds_object, stream_ptr); + auto initialize = detail::ConsumeDLPackWithStrides(initialize_object, stream_ptr); + + if (random_state_count == 0) { + return; + } + + InitializeRandomStates(reinterpret_cast(random_states.raw_data()), + random_seeds.data(), + initialize.data(), + static_cast(random_state_count), + reinterpret_cast(stream_ptr)); + }, + py::arg("random_states"), + py::arg("random_state_count"), + py::arg("random_seeds"), + py::arg("initialize"), + py::arg("stream_ptr")); + + module.def( + "verify_target_block", + [](py::handle probabilities_object, + py::handle probability_token_ids_object, + py::handle kept_count_object, + py::handle verification_draft_ids_object, + py::handle greedy_object, + py::handle logits_active_object, + py::handle random_states_object, + py::handle random_state_indices_object, + py::handle request_token_ids_ptrs_object, + py::handle entry_sequence_length_object, + py::handle request_to_generation_offsets_object, + py::handle speculative_row_object, + py::handle selected_span_ids_object, + py::handle accept_len_object, + py::handle accepted_draft_count_object, + int position_count, + uintptr_t stream_ptr) { + CudaDeviceGuard guard{GetCudaOrdinal(probabilities_object)}; + auto probabilities = detail::ConsumeDLPackWithStrides(probabilities_object, stream_ptr); + auto probability_token_ids = detail::ConsumeDLPackWithStrides(probability_token_ids_object, stream_ptr); + auto kept_count = detail::ConsumeDLPackWithStrides(kept_count_object, stream_ptr); + auto verification_draft_ids = detail::ConsumeDLPackWithStrides(verification_draft_ids_object, stream_ptr); + auto greedy = detail::ConsumeDLPackWithStrides(greedy_object, stream_ptr); + auto logits_active = detail::ConsumeDLPackWithStrides(logits_active_object, stream_ptr); + auto random_states = detail::ConsumeDLPackWithStrides(random_states_object, stream_ptr); + auto random_state_indices = detail::ConsumeDLPackWithStrides(random_state_indices_object, stream_ptr); + auto request_token_ids_ptrs = + detail::ConsumeDLPackWithStrides(request_token_ids_ptrs_object, stream_ptr); + auto entry_sequence_length = + detail::ConsumeDLPackWithStrides(entry_sequence_length_object, stream_ptr); + auto request_to_generation_offsets = + detail::ConsumeDLPackWithStrides(request_to_generation_offsets_object, stream_ptr); + auto speculative_row = detail::ConsumeDLPackWithStrides(speculative_row_object, stream_ptr); + auto selected_span_ids = detail::ConsumeDLPackWithStrides(selected_span_ids_object, stream_ptr); + auto accept_len = detail::ConsumeDLPackWithStrides(accept_len_object, stream_ptr); + auto accepted_draft_count = + accepted_draft_count_object.is_none() ? + core::Tensor{} : + detail::ConsumeDLPackWithStrides(accepted_draft_count_object, stream_ptr); + + VerifyTargetBlockParams params{}; + params.probabilities = probabilities.data_or(static_cast(nullptr)); + params.probability_stride = static_cast(probabilities.stride(0)); + params.probability_token_ids = probability_token_ids.data_or(static_cast(nullptr)); + params.token_id_stride = static_cast(probability_token_ids.stride(0)); + params.kept_count = kept_count.data_or(static_cast(nullptr)); + params.verification_draft_ids = verification_draft_ids.data_or(static_cast(nullptr)); + params.draft_row_stride = static_cast(verification_draft_ids.stride(0)); + params.greedy = greedy.data_or(static_cast(nullptr)); + params.logits_active = logits_active.data_or(static_cast(nullptr)); + params.random_states = + reinterpret_cast(random_states.data_or(static_cast(nullptr))); + params.random_state_indices = random_state_indices.data_or(static_cast(nullptr)); + params.request_token_ids_ptrs = reinterpret_cast( + request_token_ids_ptrs.data_or(static_cast(nullptr))); + params.entry_sequence_length = entry_sequence_length.data_or(static_cast(nullptr)); + params.request_to_generation_offsets = + request_to_generation_offsets.data_or(static_cast(nullptr)); + params.speculative_row = speculative_row.data_or(static_cast(nullptr)); + params.selected_span_ids = selected_span_ids.data_or(static_cast(nullptr)); + params.selected_span_stride = static_cast(selected_span_ids.stride(0)); + params.accept_len = accept_len.data_or(static_cast(nullptr)); + params.accepted_draft_count = accepted_draft_count_object.is_none() ? + nullptr : + accepted_draft_count.data_or(static_cast(nullptr)); + params.request_count = static_cast(entry_sequence_length.shape(0)); + params.generation_count = static_cast(random_state_indices.shape(0)); + params.position_count = position_count; + + invokeVerifyTargetBlock(params, reinterpret_cast(stream_ptr)); + }, + py::arg("probabilities"), + py::arg("probability_token_ids"), + py::arg("kept_count"), + py::arg("verification_draft_ids"), + py::arg("greedy"), + py::arg("logits_active"), + py::arg("random_states"), + py::arg("random_state_indices"), + py::arg("request_token_ids_ptrs"), + py::arg("entry_sequence_length"), + py::arg("request_to_generation_offsets"), + py::arg("speculative_row"), + py::arg("selected_span_ids"), + py::arg("accept_len"), + py::arg("accepted_draft_count"), + py::arg("position_count"), + py::arg("stream_ptr")); + + module.def( + "sample_processed_probabilities", + [](py::handle probabilities_object, + py::handle indices_object, + py::handle kept_object, + py::handle curand_states_object, + py::handle curand_state_indices_object, + py::handle sample_mask_object, + py::handle selected_tokens_object, + py::handle sampled_logprobs_object, + py::handle sampled_indexes_object, + py::handle sampled_nums_object, + uintptr_t stream_ptr) { + CudaDeviceGuard guard{GetCudaOrdinal(probabilities_object)}; + auto probabilities = detail::ConsumeDLPackWithStrides(probabilities_object, stream_ptr); + auto indices = detail::ConsumeDLPackWithStrides(indices_object, stream_ptr); + auto kept = detail::ConsumeDLPackWithStrides(kept_object, stream_ptr); + auto curand_states = detail::ConsumeDLPackWithStrides(curand_states_object, stream_ptr); + auto curand_state_indices = detail::ConsumeDLPackWithStrides(curand_state_indices_object, stream_ptr); + auto sample_mask = sample_mask_object.is_none() ? + core::Tensor{} : + detail::ConsumeDLPackWithStrides(sample_mask_object, stream_ptr); + auto selected_tokens = detail::ConsumeDLPackWithStrides(selected_tokens_object, stream_ptr); + auto sampled_logprobs = sampled_logprobs_object.is_none() ? + core::Tensor{} : + detail::ConsumeDLPackWithStrides(sampled_logprobs_object, stream_ptr); + auto sampled_indexes = sampled_indexes_object.is_none() ? + core::Tensor{} : + detail::ConsumeDLPackWithStrides(sampled_indexes_object, stream_ptr); + auto sampled_nums = sampled_nums_object.is_none() ? + core::Tensor{} : + detail::ConsumeDLPackWithStrides(sampled_nums_object, stream_ptr); + + SamplingParams params{}; + params.probabilities = probabilities.data_or(static_cast(nullptr)); + params.stride = static_cast(probabilities.stride(0)); + params.indices = indices.data_or(static_cast(nullptr)); + params.kept = kept.data_or(static_cast(nullptr)); + params.curandstate = + reinterpret_cast(curand_states.data_or(static_cast(nullptr))); + params.curandstate_indices = curand_state_indices.data_or(static_cast(nullptr)); + params.sample_mask = + sample_mask_object.is_none() ? nullptr : sample_mask.data_or(static_cast(nullptr)); + params.batch_size = static_cast(probabilities.shape(0)); + params.selected_tokens = selected_tokens.data_or(static_cast(nullptr)); + params.sampled_logprobs = + sampled_logprobs_object.is_none() ? nullptr : sampled_logprobs.data_or(static_cast(nullptr)); + params.sampled_indexes = + sampled_indexes_object.is_none() ? nullptr : sampled_indexes.data_or(static_cast(nullptr)); + params.sampled_nums = + sampled_nums_object.is_none() ? nullptr : sampled_nums.data_or(static_cast(nullptr)); + + invokeSampling(params, reinterpret_cast(stream_ptr)); + }, + py::arg("probabilities"), + py::arg("indices"), + py::arg("kept"), + py::arg("curand_states"), + py::arg("curand_state_indices"), + py::arg("sample_mask"), + py::arg("selected_tokens"), + py::arg("sampled_logprobs"), + py::arg("sampled_indexes"), + py::arg("sampled_nums"), + py::arg("stream_ptr")); + + module.def( + "append_one_token_and_advance_sequence", + [](py::handle token_ids_ptrs_object, + py::handle selected_tokens_object, + py::handle sequence_length_object, + uintptr_t stream_ptr) { + CudaDeviceGuard guard{GetCudaOrdinal(token_ids_ptrs_object)}; + auto token_ids_ptrs = detail::ConsumeDLPackWithStrides(token_ids_ptrs_object, stream_ptr); + auto selected_tokens = detail::ConsumeDLPackWithStrides(selected_tokens_object, stream_ptr); + auto sequence_length = detail::ConsumeDLPackWithStrides(sequence_length_object, stream_ptr); + + invokeAppendOneTokenAndAdvanceSequence( + reinterpret_cast(token_ids_ptrs.data_or(static_cast(nullptr))), + selected_tokens.data_or(static_cast(nullptr)), + sequence_length.data_or(static_cast(nullptr)), + static_cast(token_ids_ptrs.shape(0)), + reinterpret_cast(stream_ptr)); + }, + py::arg("token_ids_ptrs"), + py::arg("selected_tokens"), + py::arg("sequence_length"), + py::arg("stream_ptr")); +} + +} // namespace turbomind::python diff --git a/src/turbomind/kernels/speculative_sequence_kernels.cu b/src/turbomind/kernels/speculative_sequence_kernels.cu new file mode 100644 index 0000000000..4e6331397f --- /dev/null +++ b/src/turbomind/kernels/speculative_sequence_kernels.cu @@ -0,0 +1,489 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#include + +#include +#include +#include +#include +#include + +#include "src/turbomind/kernels/core/math.h" +#include "src/turbomind/kernels/speculative_sequence_kernels.h" + +namespace turbomind { + +__global__ void build_target_inputs(int* target_input_ids, + int* target_key_lengths, + const int* const* request_token_ids_ptrs, + const int* target_q_offsets, + const int* sequence_length, + const bool* target_ids_from_row, + const bool* finished, + int request_count) +{ + const int b = blockIdx.x; + if (b >= request_count) { + return; + } + + const int q_begin = target_q_offsets[b]; + const int q_len = target_q_offsets[b + 1] - q_begin; + const int S = sequence_length[b]; + const bool from_row = target_ids_from_row[b]; + + if (threadIdx.x == 0) { + target_key_lengths[b] = S + (from_row ? q_len - 1 : 0); + } + + if (!from_row) { + return; + } + + if (finished[b]) { + for (int j = threadIdx.x; j < q_len; j += blockDim.x) { + target_input_ids[q_begin + j] = 0; + } + return; + } + + const int token_begin = S - 1; + + for (int j = threadIdx.x; j < q_len; j += blockDim.x) { + target_input_ids[q_begin + j] = request_token_ids_ptrs[b][token_begin + j]; + } +} + +void invokeBuildTargetInputs(int* target_input_ids, + int* target_key_lengths, + const int* const* request_token_ids_ptrs, + const int* target_q_offsets, + const int* sequence_length, + const bool* target_ids_from_row, + const bool* finished, + int request_count, + cudaStream_t stream) +{ + if (request_count == 0) { + return; + } + + build_target_inputs<<>>(target_input_ids, + target_key_lengths, + request_token_ids_ptrs, + target_q_offsets, + sequence_length, + target_ids_from_row, + finished, + request_count); +} + +template +__global__ void BuildDraftExtensionKeyOffsetsKernel(int* k_offsets, + const int* q_offsets, + const int* entry_sequence_length, + const int* accept_len, + int batch_size, + int extension_index) +{ + using BlockScan = cub::BlockScan; + __shared__ typename BlockScan::TempStorage scan_storage; + + const int end = ((batch_size + BLOCK_SIZE - 1) / BLOCK_SIZE) * BLOCK_SIZE; + + int prefix = 0; + + for (int b = threadIdx.x; b < end; b += BLOCK_SIZE) { + if (b >= BLOCK_SIZE) { + __syncthreads(); + } + + int key_len = 0; + + if (b < batch_size) { + const int q_width = q_offsets[b + 1] - q_offsets[b]; + + if (q_width == 1) { + key_len = entry_sequence_length[b] + accept_len[b] + extension_index; + } + } + + int tile_sum = 0; + BlockScan{scan_storage}.ExclusiveSum(key_len, key_len, tile_sum); + + if (b < batch_size) { + k_offsets[b] = prefix + key_len; + } + + prefix += tile_sum; + } + + if (threadIdx.x == 0) { + k_offsets[batch_size] = prefix; + } +} + +void invokeBuildDraftExtensionKeyOffsets(int* k_offsets, + const int* q_offsets, + const int* entry_sequence_length, + const int* accept_len, + int batch_size, + int extension_index, + cudaStream_t stream) +{ + constexpr int block_size = 256; + + BuildDraftExtensionKeyOffsetsKernel<<<1, block_size, 0, stream>>>( + k_offsets, q_offsets, entry_sequence_length, accept_len, batch_size, extension_index); +} + +__global__ void InitializeTargetVerificationKernel(bool* block_logits_active, + int* effective_history, + int* verification_draft_ids, + const int* const* request_token_ids_ptrs, + const int* entry_sequence_length, + const bool* finished_on_entry, + const bool* speculative_row, + int* accepted_draft_count, + const int* request_to_generation_row_offsets, + int request_count, + int generation_count, + int position_count) +{ + const int b = blockIdx.x * blockDim.x + threadIdx.x; + + if (b >= request_count) { + return; + } + + if (accepted_draft_count) { + accepted_draft_count[b] = speculative_row[b] && !finished_on_entry[b] ? 0 : -1; + } + + const int g_begin = request_to_generation_row_offsets[b]; + const int g_end = request_to_generation_row_offsets[b + 1]; + + if (g_end == g_begin) { + return; + } + + const int g = g_begin; + const int S = entry_sequence_length[b]; + const bool active = !finished_on_entry[b]; + const bool verify = active && speculative_row[b]; + + for (int position = 0; position < position_count; ++position) { + const int flat = position * generation_count + g; + block_logits_active[flat] = active && (position == 0 || verify); + effective_history[flat] = S + position; + + if (position + 1 < position_count) { + verification_draft_ids[position * generation_count + g] = + verify ? request_token_ids_ptrs[b][S + position] : 0; + } + } +} + +void invokeInitializeTargetVerification(bool* block_logits_active, + int* effective_history, + int* verification_draft_ids, + const int* const* request_token_ids_ptrs, + const int* entry_sequence_length, + const bool* finished_on_entry, + const bool* speculative_row, + int* accepted_draft_count, + const int* request_to_generation_row_offsets, + int request_count, + int generation_count, + int position_count, + cudaStream_t stream) +{ + if (request_count == 0) { + return; + } + + static_cast(generation_count); + + constexpr int block_size = 128; + const int grid_size = cdiv(request_count, block_size); + + InitializeTargetVerificationKernel<<>>(block_logits_active, + effective_history, + verification_draft_ids, + request_token_ids_ptrs, + entry_sequence_length, + finished_on_entry, + speculative_row, + accepted_draft_count, + request_to_generation_row_offsets, + request_count, + generation_count, + position_count); +} + +__global__ void build_draft_refresh_inputs(int* draft_input_ids, + int* selected_token_pos, + bool* candidate_active, + const int* const* token_ids_ptrs, + const int* q_offsets, + const int* k_offsets, + const int* extension_q_offsets, + const int* accept_len, + const bool* limit_to_accept_len, + const bool* finished, + int batch_size) +{ + const int b = blockIdx.x; + + if (b >= batch_size) { + return; + } + + const int q_begin = q_offsets[b]; + const int q_end = q_offsets[b + 1]; + const int q_len = q_end - q_begin; + + if (threadIdx.x == 0) { + const int candidate_begin = extension_q_offsets[b]; + const int candidate_end = extension_q_offsets[b + 1]; + + if (candidate_end != candidate_begin) { + const int candidate = candidate_begin; + + bool active = false; + int selected_pos = 0; + + if (!finished[b]) { + if (limit_to_accept_len[b]) { + const int committed = accept_len[b]; + + if (committed > 0) { + active = true; + selected_pos = q_begin + committed - 1; + } + } + else if (q_len > 0) { + active = true; + selected_pos = q_end - 1; + } + } + + selected_token_pos[candidate] = selected_pos; + candidate_active[candidate] = active; + } + } + + if (q_len == 0) { + return; + } + + const int k_len = k_offsets[b + 1] - k_offsets[b]; + + const int token_begin = k_len - q_len + 1; + + int valid_len = q_len; + + if (limit_to_accept_len[b]) { + valid_len = accept_len[b]; + } + + for (int j = threadIdx.x; j < q_len; j += blockDim.x) { + int token = 0; + + if (j < valid_len) { + token = token_ids_ptrs[b][token_begin + j]; + } + + draft_input_ids[q_begin + j] = token; + } +} + +void invokeBuildDraftRefreshInputs(int* draft_input_ids, + int* selected_token_pos, + bool* candidate_active, + const int* const* token_ids_ptrs, + const int* refresh_q_offsets, + const int* refresh_k_offsets, + const int* extension_q_offsets, + const int* accept_len, + const bool* limit_to_accept_len, + const bool* finished, + int draft_input_count, + int batch_size, + int candidate_count, + cudaStream_t stream) +{ + if (batch_size == 0) { + return; + } + + if (draft_input_count == 0) { + return; + } + + constexpr int block_size = 256; + + build_draft_refresh_inputs<<>>(draft_input_ids, + selected_token_pos, + candidate_active, + token_ids_ptrs, + refresh_q_offsets, + refresh_k_offsets, + extension_q_offsets, + accept_len, + limit_to_accept_len, + finished, + batch_size); +} + +struct DraftArgMax { + float value; + int token_id; +}; + +struct DraftArgMaxOp { + __device__ DraftArgMax operator()(const DraftArgMax& a, const DraftArgMax& b) const + { + if (b.value > a.value) { + return b; + } + if (b.value == a.value && b.token_id < a.token_id) { + return b; + } + return a; + } +}; + +template +__global__ void DraftArgmaxAndStoreTokenKernel(const T* logits, + int logits_stride, + int vocab_size, + int* proposal_ids, + int* const* token_ids_ptrs, + const int* extension_q_offsets, + const bool* candidate_active, + const int* entry_sequence_length, + const int* accept_len, + int batch_size, + int proposal_index) +{ + const int b = blockIdx.x; + + if (b >= batch_size) { + return; + } + + const int candidate_begin = extension_q_offsets[b]; + const int candidate_end = extension_q_offsets[b + 1]; + + if (candidate_begin == candidate_end) { + return; + } + + const int c = candidate_begin; + + if (!candidate_active[c]) { + if (threadIdx.x == 0) { + proposal_ids[c] = 0; + } + return; + } + + const T* row = logits + static_cast(c) * logits_stride; + + DraftArgMax local{-CUDART_INF_F, INT_MAX}; + + for (int token_id = threadIdx.x; token_id < vocab_size; token_id += BLOCK_SIZE) { + float value = static_cast(row[token_id]); + + if (isnan(value)) { + value = -CUDART_INF_F; + } + + local = DraftArgMaxOp{}(local, DraftArgMax{value, token_id}); + } + + using BlockReduce = cub::BlockReduce; + + __shared__ typename BlockReduce::TempStorage storage; + + const DraftArgMax best = BlockReduce(storage).Reduce(local, DraftArgMaxOp{}); + + if (threadIdx.x == 0) { + proposal_ids[c] = best.token_id; + + const int token_position = entry_sequence_length[b] + accept_len[b] + proposal_index; + + token_ids_ptrs[b][token_position] = best.token_id; + } +} + +void invokeDraftArgmaxAndStoreToken(const Tensor& logits, + int* proposal_ids, + int* const* token_ids_ptrs, + const int* extension_q_offsets, + const bool* candidate_active, + const int* entry_sequence_length, + const int* accept_len, + int batch_size, + int candidate_count, + int proposal_index, + int vocab_size, + cudaStream_t stream) +{ + if (batch_size == 0 || candidate_count == 0) { + return; + } + + constexpr int block_size = 256; + const int logits_stride = static_cast(logits.shape(1)); + + auto dispatch = [&](auto t) { + using T = decltype(t); + DraftArgmaxAndStoreTokenKernel<<>>(logits.data(), + logits_stride, + vocab_size, + proposal_ids, + token_ids_ptrs, + extension_q_offsets, + candidate_active, + entry_sequence_length, + accept_len, + batch_size, + proposal_index); + }; + + TM_DISPATCH_DTYPES(logits.dtype(), dispatch, half_t, bfloat16_t); +} + +__global__ void +AdvanceSequenceByAcceptedSpan(int* sequence_length, const int* accept_len, int* accepted_draft_count, int batch_size) +{ + const int b = blockIdx.x * blockDim.x + threadIdx.x; + + if (b >= batch_size) { + return; + } + + sequence_length[b] += accept_len[b]; + + if (accepted_draft_count && accepted_draft_count[b] >= 0) { + accepted_draft_count[b] = min(accepted_draft_count[b], accept_len[b]); + } +} + +void invokeAdvanceSequenceByAcceptedSpan( + int* sequence_length, const int* accept_len, int* accepted_draft_count, int batch_size, cudaStream_t stream) +{ + if (batch_size == 0) { + return; + } + + constexpr int threads = 256; + const int blocks = (batch_size + threads - 1) / threads; + + AdvanceSequenceByAcceptedSpan<<>>( + sequence_length, accept_len, accepted_draft_count, batch_size); +} + +} // namespace turbomind diff --git a/src/turbomind/kernels/speculative_sequence_kernels.h b/src/turbomind/kernels/speculative_sequence_kernels.h new file mode 100644 index 0000000000..568b951462 --- /dev/null +++ b/src/turbomind/kernels/speculative_sequence_kernels.h @@ -0,0 +1,74 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#pragma once + +#include + +#include "src/turbomind/core/core.h" + +namespace turbomind { + +void invokeBuildTargetInputs(int* target_input_ids, + int* target_key_lengths, + const int* const* request_token_ids_ptrs, + const int* target_q_offsets, + const int* sequence_length, + const bool* target_ids_from_row, + const bool* finished, + int request_count, + cudaStream_t stream); + +void invokeBuildDraftExtensionKeyOffsets(int* k_offsets, + const int* q_offsets, + const int* entry_sequence_length, + const int* accept_len, + int batch_size, + int extension_index, + cudaStream_t stream); + +void invokeInitializeTargetVerification(bool* block_logits_active, + int* effective_history, + int* verification_draft_ids, + const int* const* request_token_ids_ptrs, + const int* entry_sequence_length, + const bool* finished_on_entry, + const bool* speculative_row, + int* accepted_draft_count, + const int* request_to_generation_row_offsets, + int request_count, + int generation_count, + int position_count, + cudaStream_t stream); + +void invokeBuildDraftRefreshInputs(int* draft_input_ids, + int* selected_token_pos, + bool* candidate_active, + const int* const* token_ids_ptrs, + const int* refresh_q_offsets, + const int* refresh_k_offsets, + const int* extension_q_offsets, + const int* accept_len, + const bool* limit_to_accept_len, + const bool* finished, + int draft_input_count, + int batch_size, + int candidate_count, + cudaStream_t stream); + +void invokeDraftArgmaxAndStoreToken(const Tensor& logits, + int* proposal_ids, + int* const* token_ids_ptrs, + const int* extension_q_offsets, + const bool* candidate_active, + const int* entry_sequence_length, + const int* accept_len, + int batch_size, + int candidate_count, + int proposal_index, + int vocab_size, + cudaStream_t stream); + +void invokeAdvanceSequenceByAcceptedSpan( + int* sequence_length, const int* accept_len, int* accepted_draft_count, int batch_size, cudaStream_t stream); + +} // namespace turbomind diff --git a/src/turbomind/kernels/speculative_sequence_python_bind.cpp b/src/turbomind/kernels/speculative_sequence_python_bind.cpp new file mode 100644 index 0000000000..88935de10b --- /dev/null +++ b/src/turbomind/kernels/speculative_sequence_python_bind.cpp @@ -0,0 +1,295 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#include + +#include + +#include + +#include "src/turbomind/kernels/speculative_sequence_kernels.h" +#include "src/turbomind/kernels/stop_criteria_kernels.h" +#include "src/turbomind/python/eagle3_component_bindings.h" +#include "src/turbomind/python/eagle3_dlpack_internal.h" +#include "src/turbomind/utils/cuda_utils.h" + +namespace py = pybind11; + +namespace turbomind::python { +namespace { + +int GetCudaOrdinal(py::handle tensor) +{ + return tensor.attr("__dlpack_device__")().cast()[1].cast(); +} + +detail::Tensor ConsumeStopWords(py::handle object, int batch_size, uintptr_t stream_ptr, int& width) +{ + if (object.is_none()) { + width = 0; + return {}; + } + + width = object.attr("shape").attr("__getitem__")(2).cast(); + py::object flattened = object.attr("view")(batch_size, 2 * width); + return detail::ConsumeDLPackWithStrides(flattened, stream_ptr); +} + +} // namespace + +void BindSpeculativeSequence(py::module_& module) +{ + module.def( + "initialize_target_verification", + [](py::handle block_logits_active_object, + py::handle effective_history_object, + py::handle verification_draft_ids_object, + py::handle request_token_ids_ptrs_object, + py::handle entry_sequence_length_object, + py::handle finished_on_entry_object, + py::handle speculative_row_object, + py::handle accepted_draft_count_object, + py::handle request_to_generation_row_offsets_object, + uintptr_t stream_ptr) { + CudaDeviceGuard guard{GetCudaOrdinal(entry_sequence_length_object)}; + auto block_logits_active = detail::ConsumeDLPackWithStrides(block_logits_active_object, stream_ptr); + auto effective_history = detail::ConsumeDLPackWithStrides(effective_history_object, stream_ptr); + auto verification_draft_ids = detail::ConsumeDLPackWithStrides(verification_draft_ids_object, stream_ptr); + auto request_token_ids_ptrs = detail::ConsumeDLPackWithStrides(request_token_ids_ptrs_object, stream_ptr); + auto entry_sequence_length = detail::ConsumeDLPackWithStrides(entry_sequence_length_object, stream_ptr); + auto finished_on_entry = detail::ConsumeDLPackWithStrides(finished_on_entry_object, stream_ptr); + auto speculative_row = detail::ConsumeDLPackWithStrides(speculative_row_object, stream_ptr); + detail::Tensor accepted_draft_count; + if (!accepted_draft_count_object.is_none()) { + accepted_draft_count = detail::ConsumeDLPackWithStrides(accepted_draft_count_object, stream_ptr); + } + auto request_to_generation_row_offsets = + detail::ConsumeDLPackWithStrides(request_to_generation_row_offsets_object, stream_ptr); + + invokeInitializeTargetVerification( + block_logits_active.data_or(static_cast(nullptr)), + effective_history.data_or(static_cast(nullptr)), + verification_draft_ids.data_or(static_cast(nullptr)), + reinterpret_cast(request_token_ids_ptrs.data_or(static_cast(nullptr))), + entry_sequence_length.data_or(static_cast(nullptr)), + finished_on_entry.data_or(static_cast(nullptr)), + speculative_row.data_or(static_cast(nullptr)), + accepted_draft_count.data_or(static_cast(nullptr)), + request_to_generation_row_offsets.data_or(static_cast(nullptr)), + static_cast(entry_sequence_length.shape(0)), + static_cast(block_logits_active.shape(1)), + static_cast(block_logits_active.shape(0)), + reinterpret_cast(stream_ptr)); + }, + py::arg("block_logits_active"), + py::arg("effective_history"), + py::arg("verification_draft_ids"), + py::arg("request_token_ids_ptrs"), + py::arg("entry_sequence_length"), + py::arg("finished_on_entry"), + py::arg("speculative_row"), + py::arg("accepted_draft_count"), + py::arg("request_to_generation_row_offsets"), + py::arg("stream_ptr")); + + module.def( + "build_draft_extension_key_offsets", + [](py::handle k_offsets_object, + py::handle q_offsets_object, + py::handle entry_sequence_length_object, + py::handle accept_len_object, + int extension_index, + uintptr_t stream_ptr) { + CudaDeviceGuard guard{GetCudaOrdinal(q_offsets_object)}; + auto k_offsets = detail::ConsumeDLPackWithStrides(k_offsets_object, stream_ptr); + auto q_offsets = detail::ConsumeDLPackWithStrides(q_offsets_object, stream_ptr); + auto entry_sequence_length = detail::ConsumeDLPackWithStrides(entry_sequence_length_object, stream_ptr); + auto accept_len = detail::ConsumeDLPackWithStrides(accept_len_object, stream_ptr); + + invokeBuildDraftExtensionKeyOffsets(k_offsets.data_or(static_cast(nullptr)), + q_offsets.data_or(static_cast(nullptr)), + entry_sequence_length.data_or(static_cast(nullptr)), + accept_len.data_or(static_cast(nullptr)), + static_cast(entry_sequence_length.shape(0)), + extension_index, + reinterpret_cast(stream_ptr)); + }, + py::arg("k_offsets"), + py::arg("q_offsets"), + py::arg("entry_sequence_length"), + py::arg("accept_len"), + py::arg("extension_index"), + py::arg("stream_ptr")); + + module.def( + "build_draft_refresh_inputs", + [](py::handle draft_input_ids_object, + py::handle selected_token_pos_object, + py::handle candidate_active_object, + py::handle token_ids_ptrs_object, + py::handle refresh_q_offsets_object, + py::handle refresh_k_offsets_object, + py::handle extension_q_offsets_object, + py::handle accept_len_object, + py::handle limit_to_accept_len_object, + py::handle finished_object, + uintptr_t stream_ptr) { + CudaDeviceGuard guard{GetCudaOrdinal(token_ids_ptrs_object)}; + auto draft_input_ids = detail::ConsumeDLPackWithStrides(draft_input_ids_object, stream_ptr); + auto selected_token_pos = detail::ConsumeDLPackWithStrides(selected_token_pos_object, stream_ptr); + auto candidate_active = detail::ConsumeDLPackWithStrides(candidate_active_object, stream_ptr); + auto token_ids_ptrs = detail::ConsumeDLPackWithStrides(token_ids_ptrs_object, stream_ptr); + auto refresh_q_offsets = detail::ConsumeDLPackWithStrides(refresh_q_offsets_object, stream_ptr); + auto refresh_k_offsets = detail::ConsumeDLPackWithStrides(refresh_k_offsets_object, stream_ptr); + auto extension_q_offsets = detail::ConsumeDLPackWithStrides(extension_q_offsets_object, stream_ptr); + auto accept_len = detail::ConsumeDLPackWithStrides(accept_len_object, stream_ptr); + auto limit_to_accept_len = detail::ConsumeDLPackWithStrides(limit_to_accept_len_object, stream_ptr); + auto finished = detail::ConsumeDLPackWithStrides(finished_object, stream_ptr); + + invokeBuildDraftRefreshInputs( + draft_input_ids.data_or(static_cast(nullptr)), + selected_token_pos.data_or(static_cast(nullptr)), + candidate_active.data_or(static_cast(nullptr)), + reinterpret_cast(token_ids_ptrs.data_or(static_cast(nullptr))), + refresh_q_offsets.data_or(static_cast(nullptr)), + refresh_k_offsets.data_or(static_cast(nullptr)), + extension_q_offsets.data_or(static_cast(nullptr)), + accept_len.data_or(static_cast(nullptr)), + limit_to_accept_len.data_or(static_cast(nullptr)), + finished.data_or(static_cast(nullptr)), + static_cast(draft_input_ids.shape(0)), + static_cast(accept_len.shape(0)), + static_cast(selected_token_pos.shape(0)), + reinterpret_cast(stream_ptr)); + }, + py::arg("draft_input_ids"), + py::arg("selected_token_pos"), + py::arg("candidate_active"), + py::arg("token_ids_ptrs"), + py::arg("refresh_q_offsets"), + py::arg("refresh_k_offsets"), + py::arg("extension_q_offsets"), + py::arg("accept_len"), + py::arg("limit_to_accept_len"), + py::arg("finished"), + py::arg("stream_ptr")); + + module.def( + "draft_argmax_and_store_token", + [](py::handle logits_object, + py::handle proposal_ids_object, + py::handle token_ids_ptrs_object, + py::handle extension_q_offsets_object, + py::handle candidate_active_object, + py::handle entry_sequence_length_object, + py::handle accept_len_object, + int proposal_index, + int vocab_size, + uintptr_t stream_ptr) { + CudaDeviceGuard guard{GetCudaOrdinal(logits_object)}; + auto logits = detail::ConsumeDLPackWithStrides(logits_object, stream_ptr); + auto proposal_ids = detail::ConsumeDLPackWithStrides(proposal_ids_object, stream_ptr); + auto token_ids_ptrs = detail::ConsumeDLPackWithStrides(token_ids_ptrs_object, stream_ptr); + auto extension_q_offsets = detail::ConsumeDLPackWithStrides(extension_q_offsets_object, stream_ptr); + auto candidate_active = detail::ConsumeDLPackWithStrides(candidate_active_object, stream_ptr); + auto entry_sequence_length = detail::ConsumeDLPackWithStrides(entry_sequence_length_object, stream_ptr); + auto accept_len = detail::ConsumeDLPackWithStrides(accept_len_object, stream_ptr); + + invokeDraftArgmaxAndStoreToken( + logits, + proposal_ids.data_or(static_cast(nullptr)), + reinterpret_cast(token_ids_ptrs.data_or(static_cast(nullptr))), + extension_q_offsets.data_or(static_cast(nullptr)), + candidate_active.data_or(static_cast(nullptr)), + entry_sequence_length.data_or(static_cast(nullptr)), + accept_len.data_or(static_cast(nullptr)), + static_cast(token_ids_ptrs.shape(0)), + static_cast(logits.shape(0)), + proposal_index, + vocab_size, + reinterpret_cast(stream_ptr)); + }, + py::arg("logits"), + py::arg("proposal_ids"), + py::arg("token_ids_ptrs"), + py::arg("extension_q_offsets"), + py::arg("candidate_active"), + py::arg("entry_sequence_length"), + py::arg("accept_len"), + py::arg("proposal_index"), + py::arg("vocab_size"), + py::arg("stream_ptr")); + + module.def( + "stop_criteria", + [](py::handle token_ids_ptrs_object, + py::handle sequence_length_object, + py::handle stop_words_object, + py::handle sequence_length_limit_object, + py::handle finished_object, + uintptr_t stream_ptr) { + CudaDeviceGuard guard{GetCudaOrdinal(token_ids_ptrs_object)}; + auto token_ids_ptrs = detail::ConsumeDLPackWithStrides(token_ids_ptrs_object, stream_ptr); + auto sequence_length = detail::ConsumeDLPackWithStrides(sequence_length_object, stream_ptr); + int stop_words_width = 0; + auto stop_words = ConsumeStopWords( + stop_words_object, static_cast(sequence_length.shape(0)), stream_ptr, stop_words_width); + auto sequence_length_limit = detail::ConsumeDLPackWithStrides(sequence_length_limit_object, stream_ptr); + auto finished = detail::ConsumeDLPackWithStrides(finished_object, stream_ptr); + + invokeStopCriteria( + reinterpret_cast(token_ids_ptrs.data_or(static_cast(nullptr))), + sequence_length.data_or(static_cast(nullptr)), + stop_words.data_or(static_cast(nullptr)), + stop_words_width, + sequence_length_limit.data_or(static_cast(nullptr)), + finished.data_or(static_cast(nullptr)), + static_cast(sequence_length.shape(0)), + reinterpret_cast(stream_ptr)); + }, + py::arg("token_ids_ptrs"), + py::arg("sequence_length"), + py::arg("stop_words"), + py::arg("sequence_length_limit"), + py::arg("finished"), + py::arg("stream_ptr")); + + module.def( + "speculative_stop_criteria", + [](py::handle token_ids_ptrs_object, + py::handle entry_sequence_length_object, + py::handle accept_len_object, + py::handle stop_words_object, + py::handle sequence_length_limit_object, + py::handle finished_object, + uintptr_t stream_ptr) { + CudaDeviceGuard guard{GetCudaOrdinal(token_ids_ptrs_object)}; + auto token_ids_ptrs = detail::ConsumeDLPackWithStrides(token_ids_ptrs_object, stream_ptr); + auto entry_sequence_length = detail::ConsumeDLPackWithStrides(entry_sequence_length_object, stream_ptr); + auto accept_len = detail::ConsumeDLPackWithStrides(accept_len_object, stream_ptr); + int stop_words_width = 0; + auto stop_words = ConsumeStopWords( + stop_words_object, static_cast(entry_sequence_length.shape(0)), stream_ptr, stop_words_width); + auto sequence_length_limit = detail::ConsumeDLPackWithStrides(sequence_length_limit_object, stream_ptr); + auto finished = detail::ConsumeDLPackWithStrides(finished_object, stream_ptr); + + invokeStopCriteria( + reinterpret_cast(token_ids_ptrs.data_or(static_cast(nullptr))), + entry_sequence_length.data_or(static_cast(nullptr)), + accept_len.data_or(static_cast(nullptr)), + stop_words.data_or(static_cast(nullptr)), + stop_words_width, + sequence_length_limit.data_or(static_cast(nullptr)), + finished.data_or(static_cast(nullptr)), + static_cast(entry_sequence_length.shape(0)), + reinterpret_cast(stream_ptr)); + }, + py::arg("token_ids_ptrs"), + py::arg("entry_sequence_length"), + py::arg("accept_len"), + py::arg("stop_words"), + py::arg("sequence_length_limit"), + py::arg("finished"), + py::arg("stream_ptr")); +} + +} // namespace turbomind::python diff --git a/src/turbomind/kernels/stop_criteria_kernels.cu b/src/turbomind/kernels/stop_criteria_kernels.cu index 424b3a933d..5dfad74056 100644 --- a/src/turbomind/kernels/stop_criteria_kernels.cu +++ b/src/turbomind/kernels/stop_criteria_kernels.cu @@ -16,91 +16,156 @@ #include "src/turbomind/kernels/core/math.h" #include "src/turbomind/kernels/stop_criteria_kernels.h" -#include "src/turbomind/utils/cuda_utils.h" namespace turbomind { -__global__ void stop_words_criterion_v2(const int** token_ids_ptrs, - const int* sequence_length, - const int* stop_words, - bool* finished, - int stop_words_len, - int batch_size) +template +__global__ void stop_criteria(const int* const* token_ids_ptrs, + const int* sequence_length, + int* accept_len, + const int* stop_words, + int stop_words_width, + const int* sequence_length_limit, + bool* finished, + int batch_size) { - const int id = blockIdx.x * blockDim.x + threadIdx.x; - const int batch_idx = blockIdx.y; + const int b = blockIdx.x * blockDim.x + threadIdx.x; - const int* base_stop_words = stop_words + batch_idx * 2 * stop_words_len; - const int* base_offsets = base_stop_words + stop_words_len; - - if (id >= stop_words_len || base_offsets[id] < 0) { + if (b >= batch_size) { return; } - const int item_end = base_offsets[id]; - const int item_start = (id > 0) ? base_offsets[id - 1] : 0; - const int item_size = item_end - item_start; + int entry_len = 0; + int span_len = 0; - const int seq_len = sequence_length[batch_idx]; - const int* token_ids = token_ids_ptrs[batch_idx]; + if constexpr (kSpeculative) { + if (finished[b]) { + accept_len[b] = 0; + return; + } - /* Enough previously generated tokens to look for a match */ - if (seq_len >= item_size) { - // token_ids[seq_len - 1] is the last token - for (int token_idx = item_size - 1, offset = seq_len - 1; token_idx >= 0; token_idx--, offset--) { - if (token_ids[offset] != base_stop_words[item_start + token_idx]) { - return; - } + entry_len = sequence_length[b]; + span_len = accept_len[b]; + } + else { + if (finished[b]) { + return; } - finished[batch_idx] = true; + + const int current_len = sequence_length[b]; + + if (current_len <= 0) { + return; + } + + entry_len = current_len - 1; + span_len = 1; } -} -void invokeStopWordsCriterion_v2(const int** token_ids_ptrs, - const int* sequence_length, - const int* stop_words, - bool* finished, - int stop_words_len, - int batch_size, - cudaStream_t stream) -{ - // Check if we have sampled a word from the stop_words list. If so, stop the sequence. + if (span_len <= 0) { + return; + } + + const int* tokens = token_ids_ptrs[b]; + + for (int j = 0; j < span_len; ++j) { + const int effective_len = entry_len + j + 1; + + bool terminal = effective_len >= sequence_length_limit[b]; + + if (!terminal && stop_words != nullptr) { + const int* words = stop_words + b * 2 * stop_words_width; + const int* offsets = words + stop_words_width; + + for (int phrase = 0; phrase < stop_words_width; ++phrase) { + const int phrase_end = offsets[phrase]; - const int block = std::min(round_up(stop_words_len, 32), 256); - const dim3 grid(cdiv(stop_words_len, block), batch_size); + if (phrase_end < 0) { + break; + } - stop_words_criterion_v2<<>>( - token_ids_ptrs, sequence_length, stop_words, finished, stop_words_len, batch_size); - TM_CUDA_CHECK(cudaGetLastError()); + const int phrase_begin = phrase == 0 ? 0 : offsets[phrase - 1]; + const int phrase_size = phrase_end - phrase_begin; + + if (phrase_size <= 0 || effective_len < phrase_size) { + continue; + } + + const int history_begin = effective_len - phrase_size; + bool match = true; + + for (int t = 0; t < phrase_size; ++t) { + if (tokens[history_begin + t] != words[phrase_begin + t]) { + match = false; + break; + } + } + + if (match) { + terminal = true; + break; + } + } + } + + if (terminal) { + if constexpr (kSpeculative) { + accept_len[b] = j + 1; + } + + finished[b] = true; + return; + } + } } -__global__ void length_criterion_v2(bool* finished, // - const int* sequence_length, - const int* sequence_length_limit, - int batch_size) +void invokeStopCriteria(const int* const* token_ids_ptrs, + const int* sequence_length, + const int* stop_words, + int stop_words_width, + const int* sequence_length_limit, + bool* finished, + int batch_size, + cudaStream_t stream) { - const int idx = threadIdx.x + blockDim.x * blockIdx.x; - if (idx >= batch_size) { + if (batch_size == 0) { return; } - if (sequence_length[idx] >= sequence_length_limit[idx]) { - finished[idx] = true; - } + + constexpr int block_size = 128; + stop_criteria<<>>(token_ids_ptrs, + sequence_length, + nullptr, + stop_words, + stop_words_width, + sequence_length_limit, + finished, + batch_size); } -void invokeLengthCriterion_v2(bool* finished, // - const int* sequence_length, - const int* sequence_length_limit, - int batch_size, - cudaStream_t stream) +void invokeStopCriteria(const int* const* token_ids_ptrs, + const int* entry_sequence_length, + int* accept_len, + const int* stop_words, + int stop_words_width, + const int* sequence_length_limit, + bool* finished, + int batch_size, + cudaStream_t stream) { - // Check if we have attained the sequence length limit. If so, stop the sequence. - - constexpr int block = 256; - const int grid = cdiv(batch_size, block); + if (batch_size == 0) { + return; + } - length_criterion_v2<<>>(finished, sequence_length, sequence_length_limit, batch_size); - TM_CUDA_CHECK(cudaGetLastError()); + constexpr int block_size = 128; + stop_criteria<<>>(token_ids_ptrs, + entry_sequence_length, + accept_len, + stop_words, + stop_words_width, + sequence_length_limit, + finished, + batch_size); } } // namespace turbomind diff --git a/src/turbomind/kernels/stop_criteria_kernels.h b/src/turbomind/kernels/stop_criteria_kernels.h index 41bc81ba6b..05d64308a3 100644 --- a/src/turbomind/kernels/stop_criteria_kernels.h +++ b/src/turbomind/kernels/stop_criteria_kernels.h @@ -15,24 +15,27 @@ */ #pragma once -#include - #include namespace turbomind { -void invokeStopWordsCriterion_v2(const int** token_ids_ptrs, - const int* sequence_length, - const int* stop_words, - bool* finished, - int stop_words_len, - int batch_size, - cudaStream_t stream); +void invokeStopCriteria(const int* const* token_ids_ptrs, + const int* sequence_length, + const int* stop_words, + int stop_words_width, + const int* sequence_length_limit, + bool* finished, + int batch_size, + cudaStream_t stream); -void invokeLengthCriterion_v2(bool* finished, // - const int* sequence_length, - const int* sequence_length_limit, - int batch_size, - cudaStream_t stream); +void invokeStopCriteria(const int* const* token_ids_ptrs, + const int* entry_sequence_length, + int* accept_len, + const int* stop_words, + int stop_words_width, + const int* sequence_length_limit, + bool* finished, + int batch_size, + cudaStream_t stream); } // namespace turbomind diff --git a/src/turbomind/models/CMakeLists.txt b/src/turbomind/models/CMakeLists.txt index 912ff2095e..385ab0bd7d 100644 --- a/src/turbomind/models/CMakeLists.txt +++ b/src/turbomind/models/CMakeLists.txt @@ -1,9 +1,27 @@ cmake_minimum_required(VERSION 3.25) +add_library(target_hidden_projection_kernels STATIC + speculative/eagle3/target_hidden_projection_kernels.cu) +set_property(TARGET target_hidden_projection_kernels PROPERTY POSITION_INDEPENDENT_CODE ON) +set_property(TARGET target_hidden_projection_kernels PROPERTY CUDA_RESOLVE_DEVICE_SYMBOLS ON) + add_library(models STATIC + batch_status.cc language_model.cc input_processor.cc + input_processor_speculative.cc output_processor.cc + speculative/registry.cc + speculative/fixed_chain_policy.cc + speculative/fixed_chain_setup.cc + speculative/fixed_chain_model.cc + speculative/collect_hidden_states.cc + speculative/eagle3/eagle3_model.cc + speculative/eagle3/target_hidden_projection.cc + speculative/eagle3/eagle3_weight.cc + speculative/qwen3_5_mtp/qwen3_5_mtp_weight.cc + speculative/qwen3_5_mtp/target_final_hidden.cc + speculative/qwen3_5_mtp/qwen3_5_mtp_model.cc linear_weight.cc norm_weight.cc layer_norm_weight.cc @@ -57,6 +75,11 @@ target_link_libraries(models PUBLIC memory_utils cuda_utils anomaly_handler) +target_link_libraries(models PRIVATE + target_hidden_projection_kernels + draft_carry_kernels + speculative_sequence_kernels + device_comm) target_compile_options(models PRIVATE $<$:-Xptxas=-v --generate-line-info --threads=${NVCC_THREADS}>) diff --git a/src/turbomind/models/batch_status.cc b/src/turbomind/models/batch_status.cc new file mode 100644 index 0000000000..a5133273d7 --- /dev/null +++ b/src/turbomind/models/batch_status.cc @@ -0,0 +1,172 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#include "src/turbomind/models/batch_status.h" + +#include + +#include "src/turbomind/core/copy.h" + +namespace turbomind { + +using core::BatchCopy; + +struct BatchStatus::Data { + Buffer_ sequence_length; + Buffer_ readonly_block_num; + Buffer_ finished; + + Buffer_ autoregres; + Buffer_ generating; + + int n_generating{}; + int verification_positions{}; +}; + +BatchStatus::BatchStatus(int max_batch_size, int phases): max_batch_size_{max_batch_size} +{ + false_ = {max_batch_size, kDEVICE}; + Clear(false_); + + finished_buf_ = {max_batch_size, kCPUpinned}; + finished_ = {{max_batch_size}, kBool, kDEVICE}; + + sequence_length_buf_ = {max_batch_size, kCPUpinned}; + readonly_block_num_buf_ = {max_batch_size, kCPUpinned}; + sequence_length_ = {{max_batch_size}, kInt, kDEVICE}; + + data_.reserve(phases); + for (int i = 0; i < phases; ++i) { + auto d = std::make_unique(); + d->sequence_length = empty_like(sequence_length_buf_, kDEVICE); + d->readonly_block_num = empty_like(readonly_block_num_buf_, kDEVICE); + d->finished = empty_like(finished_buf_, kDEVICE); + d->autoregres = {max_batch_size, kCPU}; + d->generating = {max_batch_size, kCPU}; + data_.push_back(std::move(d)); + } +} + +BatchStatus::~BatchStatus() = default; + +void BatchStatus::Run(BatchOp op, int phase, TensorMap& env) +{ + switch (op) { + case BatchOp::kSetup: + return Setup(phase, env); + case BatchOp::kPrepare: + return Prepare(phase, env); + case BatchOp::kUnprep: + return Unprep(phase, env); + case BatchOp::kFetch: + return Fetch(phase, env); + default: + return; + } +} + +void BatchStatus::Setup(int phase, TensorMap& env) +{ + auto& d = *data_.at(phase); + auto& copy = *env.at("copy").data()[0]; + + Buffer_ requests = env.at("requests").buffer(); + + d.n_generating = 0; + d.verification_positions = 0; + + for (int i = 0; i < requests.size(); ++i) { + const Sequence& request = *requests[i]; + const SubmittedRow& row = *request.submitted; + + d.autoregres[i] = row.autoregres; + d.generating[i] = row.generating; + d.n_generating += row.generating; + + if (row.generating) { + d.verification_positions = std::max(d.verification_positions, row.verification_positions); + } + + sequence_length_buf_[i] = row.autoregres ? request.seq_len : row.key_capacity_end; + readonly_block_num_buf_[i] = request.readonly_block_num; + } + + copy(sequence_length_buf_, requests.size(), d.sequence_length); + copy(readonly_block_num_buf_, requests.size(), d.readonly_block_num); + + env.produce("verification_positions", Buffer_{&d.verification_positions, 1, kCPU}); +} + +void BatchStatus::Prepare(int phase, TensorMap& env) +{ + auto& d = *data_.at(phase); + auto& batch = *env.at("batch").data()[0]; + auto& copy = *env.at("copy").data()[0]; + + if (auto group = copy.group()) { + for (int i = 0; i < batch.bsz; ++i) { + if (const int j = batch.perm[i]; j < batch.bs0) { + copy(finished_.front().data() + j, 1, finished_.back().data() + i); + } + else { + copy(false_.data() + i, 1, finished_.back().data() + i); + } + } + finished_.Swap(); + } + + if (auto group = copy.group()) { + for (int i = 0; i < batch.bsz; ++i) { + if (const int j = batch.perm[i]; j < batch.bs0 && d.autoregres[i]) { + copy(sequence_length_.front().data() + j, 1, sequence_length_.back().data() + i); + } + else { + copy(d.sequence_length.data() + i, 1, sequence_length_.back().data() + i); + } + } + sequence_length_.Swap(); + } + + env.produce("finished", finished_.front()); + env.produce("sequence_length", sequence_length_.front()); + env.produce("readonly_block_num", d.readonly_block_num); +} + +void BatchStatus::Unprep(int phase, TensorMap& env) +{ + auto& d = *data_.at(phase); + auto& copy = *env.at("copy").data()[0]; + + copy(sequence_length_.front().buffer(), d.sequence_length.size(), d.sequence_length); + copy(finished_.front().buffer(), d.finished.size(), d.finished); +} + +void BatchStatus::Fetch(int phase, TensorMap& env) +{ + auto& d = *data_.at(phase); + auto& copy = *env.at("copy").data()[0]; + + copy(d.sequence_length, d.sequence_length.size(), sequence_length_buf_); + env.produce("sequence_length", sequence_length_buf_); + + copy(d.finished, d.finished.size(), finished_buf_); + env.produce("finished", finished_buf_); + + env.produce("generating", d.generating); +} + +int BatchStatus::VerificationPositions(int phase) const +{ + return data_.at(phase)->verification_positions; +} + +int BatchStatus::GeneratingCount(int phase) const +{ + return data_.at(phase)->n_generating; +} + +Buffer_ BatchStatus::SequenceLength() const +{ + return sequence_length_.data_[0].buffer(); +} + +} // namespace turbomind diff --git a/src/turbomind/models/batch_status.h b/src/turbomind/models/batch_status.h new file mode 100644 index 0000000000..c08979adf2 --- /dev/null +++ b/src/turbomind/models/batch_status.h @@ -0,0 +1,49 @@ +// Copyright (c) OpenMMLab. All rights reserved. +#pragma once + +#include "src/turbomind/core/core.h" +#include "src/turbomind/core/state.h" +#include "src/turbomind/engine/batch.h" +#include "src/turbomind/engine/request.h" + +#include +#include + +namespace turbomind { + +class BatchStatus { +public: + BatchStatus(int max_batch_size, int phases); + + ~BatchStatus(); + + void Run(BatchOp op, int phase, TensorMap& env); + + int VerificationPositions(int phase) const; + + int GeneratingCount(int phase) const; + + Buffer_ SequenceLength() const; + +private: + struct Data; + + void Setup(int phase, TensorMap& env); + void Prepare(int phase, TensorMap& env); + void Unprep(int phase, TensorMap& env); + void Fetch(int phase, TensorMap& env); + + const int max_batch_size_; + + Buffer_ false_; + State finished_; + State sequence_length_; + + Buffer_ sequence_length_buf_; + Buffer_ readonly_block_num_buf_; + Buffer_ finished_buf_; + + std::vector> data_; +}; + +} // namespace turbomind diff --git a/src/turbomind/models/input_processor.cc b/src/turbomind/models/input_processor.cc index bfcee37dea..b1129b0df3 100644 --- a/src/turbomind/models/input_processor.cc +++ b/src/turbomind/models/input_processor.cc @@ -1,3 +1,4 @@ +// Copyright (c) OpenMMLab. All rights reserved. #include "src/turbomind/models/input_processor.h" @@ -6,261 +7,295 @@ #include "src/turbomind/engine/request.h" +#include "src/turbomind/models/input_processor_impl.h" #include "src/turbomind/models/vision_model.h" namespace turbomind { using std::vector; -struct InputProcessor::Impl { -public: - Impl(const EngineParam& engine, int hidden_units, DataType data_type, int phases): - max_batch_size_{engine.max_batch_size}, max_forward_token_num_{engine.max_forward_token_num} - { - input_ids_buf_ = {max_forward_token_num_, kCPUpinned}; - input_ids_offsets_buf_ = {max_batch_size_ + 1, kCPUpinned}; - decode_token_pos_buf_ = {max_batch_size_, kCPUpinned}; - - data_.reserve(phases); - for (int i = 0; i < phases; ++i) { - auto& d = data_.emplace_back(); - d.input_ids = empty_like(input_ids_buf_, kDEVICE); - d.input_ids_offsets = empty_like(input_ids_offsets_buf_, kDEVICE); - d.selected_token_pos = empty_like(decode_token_pos_buf_, kDEVICE); - - d.autoreg_ids_pos = {max_batch_size_, kCPU}; // ! CPU buffer - - /// TODO: initialize only when required - d.input_embeds_buf = {{max_forward_token_num_, hidden_units}, data_type, kCPUpinned}; +InputProcessor::Impl::Impl(const EngineParam& engine, + int hidden_units, + DataType data_type, + int phases, + bool speculative, + int max_verification_positions, + bool successor_embeddings): + max_batch_size_{engine.max_batch_size}, + max_forward_token_num_{engine.max_forward_token_num}, + speculative_engine_{speculative}, + successor_embeddings_{successor_embeddings} +{ + input_ids_buf_ = {max_forward_token_num_, kCPUpinned}; + input_ids_offsets_buf_ = {max_batch_size_ + 1, kCPUpinned}; + decode_token_pos_buf_ = {max_batch_size_ * max_verification_positions, kCPUpinned}; + target_ids_from_row_buf_ = {max_batch_size_, kCPUpinned}; + + data_.reserve(phases); + for (int i = 0; i < phases; ++i) { + auto& d = data_.emplace_back(); + d.input_ids = empty_like(input_ids_buf_, kDEVICE); + d.input_ids_offsets = empty_like(input_ids_offsets_buf_, kDEVICE); + d.selected_token_pos = empty_like(decode_token_pos_buf_, kDEVICE); + d.target_ids_from_row = empty_like(target_ids_from_row_buf_, kDEVICE); + if (speculative_engine_) { + d.target_key_lengths = {max_batch_size_, kDEVICE}; } - } - int Add(Sequence& c) - { - const auto& r = *c.req; - - // trim input embeds - if (!c.input_embeds_offsets.empty()) { - Interval l{0, (int)c.tokens.size()}; - using Size = Interval::Size; - auto& embeds = c.input_embeds; - auto& offsets = c.input_embeds_offsets; - int i = embeds.size() - 1; - for (; i >= 0; --i) { - Interval r{offsets[i], Size{(int)embeds[i].shape(0)}}; - if (auto o = r & l) { - if (o.end() < r.end()) { - embeds[i] = embeds[i].slice(0, o.end() - r.begin()); - } - break; - } - } - embeds.resize(i + 1); - offsets.resize(i + 1); - } + d.autoreg_ids_pos = {max_batch_size_, kCPU}; // ! CPU buffer - if (auto ranges_ptr = r.inputs.try_("input_embedding_ranges")) { // [n, 2] - auto embeds = r.inputs.at("input_embeddings"); // [k, d] - if (ranges_ptr->ndim() != 2 || embeds.ndim() != 2 || ranges_ptr->shape(1) != 2) { - /// TODO: reject for invalid shapes - return Request::kInvalid; - } + /// TODO: initialize only when required + d.input_embeds_buf = { + {max_forward_token_num_ + max_batch_size_, hidden_units}, data_type, kCPUpinned}; + } +} - const auto [sum, dim] = embeds.shapes(0, 1); - const auto n = ranges_ptr->shape(0); - const auto ranges = ranges_ptr->data(); - - int offset = 0; - int last = c.step0; - for (int i = 0; i < n; ++i) { - Interval range{c.step0 + ranges[i * 2], c.step0 + ranges[i * 2 + 1]}; - auto size = (int)range.size(); - if (range.begin() < last) { - /// TODO: reject for non-sorted ranges - return Request::kInvalid; - } - if (range.end() > c.seq_len) { - /// TODO: reject for dst range OOB - return Request::kInvalid; - } - if (offset + size > sum) { - /// TODO: reject for src range OOB - return Request::kInvalid; +int InputProcessor::Impl::Add(Sequence& c) +{ + const auto& r = *c.req; + + // trim input embeds + if (!c.input_embeds_offsets.empty()) { + Interval l{0, (int)c.tokens.size()}; + using Size = Interval::Size; + auto& embeds = c.input_embeds; + auto& offsets = c.input_embeds_offsets; + int i = embeds.size() - 1; + for (; i >= 0; --i) { + Interval r{offsets[i], Size{(int)embeds[i].shape(0)}}; + if (auto o = r & l) { + if (o.end() < r.end()) { + embeds[i] = embeds[i].slice(0, o.end() - r.begin()); } - c.input_embeds_offsets.push_back(range.begin()); - c.input_embeds.push_back(embeds.slice(offset, size)); // reference into `embeds` - offset += size; - last = range.end(); + break; } } - - return 0; + embeds.resize(i + 1); + offsets.resize(i + 1); } - void Add(int phase, TensorMap& env) - { - const Buffer_ rc = env.at("requests").buffer(); - for (int i = 0; i < rc.size(); ++i) { - auto& c = *TM_CHECK_NOTNULL(rc[i]); - if (c.status == 0) { - c.status = Add(c); - } + if (auto ranges_ptr = r.inputs.try_("input_embedding_ranges")) { // [n, 2] + auto embeds = r.inputs.at("input_embeddings"); // [k, d] + if (ranges_ptr->ndim() != 2 || embeds.ndim() != 2 || ranges_ptr->shape(1) != 2) { + /// TODO: reject for invalid shapes + return Request::kInvalid; } - } - void Setup(int phase, TensorMap& env) - { - auto& d = data_.at(phase); - auto& b = *env.at("batch").data()[0]; - auto& copy = *env.at("copy").data()[0]; - - Buffer_ rc = env.at("requests").buffer(); - - input_ids_offsets_buf_[0] = 0; - for (int i = 0; i < rc.size(); ++i) { - input_ids_offsets_buf_[i + 1] = input_ids_offsets_buf_[i]; - if (const auto& c = *rc[i]; TM_UNLIKELY(!c.autoregres)) { - const auto src = c.token_ids + c.history_len + c.inflight_input_len; - std::copy_n(src, c.input_len, input_ids_buf_.data() + input_ids_offsets_buf_[i]); - // dbg(std::vector(src, src + c.input_len)); - d.autoreg_ids_pos[i] = -1; - input_ids_offsets_buf_[i + 1] += c.input_len; + const auto [sum, dim] = embeds.shapes(0, 1); + const auto n = ranges_ptr->shape(0); + const auto ranges = ranges_ptr->data(); + + int offset = 0; + int last = c.step0; + for (int i = 0; i < n; ++i) { + Interval range{c.step0 + ranges[i * 2], c.step0 + ranges[i * 2 + 1]}; + auto size = (int)range.size(); + if (range.begin() < last) { + /// TODO: reject for non-sorted ranges + return Request::kInvalid; } - else { - d.autoreg_ids_pos[i] = input_ids_offsets_buf_[i]; - input_ids_offsets_buf_[i + 1] += 1; + if (range.end() > c.seq_len) { + /// TODO: reject for dst range OOB + return Request::kInvalid; } - decode_token_pos_buf_[i] = input_ids_offsets_buf_[i + 1] - 1; - } - - // dbg(core::to_vector(input_ids_offsets_buf_.slice(0, bsz + 1))); - // dbg(core::to_vector(decode_token_pos_buf_.slice(0, bsz))); - - copy(input_ids_buf_, input_ids_offsets_buf_[b.bsz], d.input_ids); - copy(decode_token_pos_buf_, b.bsz, d.selected_token_pos); - copy(input_ids_offsets_buf_, b.bsz + 1, d.input_ids_offsets); - - // dbg(decode_token_pos_buf_[0]); - - d.input_token_num = input_ids_offsets_buf_[b.bsz]; - // dbg(d.input_token_num); - - env.produce("token_num", Buffer{&d.input_token_num, 1, kCPU}); - - //////////////////////////////////////////////////////////////// - /// input embeddings - d.input_embeds_coords.clear(); - auto embed_ptr = (uint8_t*)d.input_embeds_buf.raw_data(); - for (int k = 0; k < rc.size(); ++k) { - if (auto& c = *rc[k]; !c.autoregres) { - const auto& embeds = c.input_embeds; - const auto& offsets = c.input_embeds_offsets; - Interval p{input_ids_offsets_buf_[k], input_ids_offsets_buf_[k + 1]}; - Interval s{c.history_len + c.inflight_input_len, p.size()}; - for (int i = (int)offsets.size() - 1; i >= 0; --i) { - Interval r{offsets[i], Interval::Size{(int)embeds[i].shape(0)}}; - auto o = r & s; - if (auto size = (int)o.size()) { - auto src = embeds[i].slice(o.begin() - r.begin(), size); - embed_ptr = std::copy_n((const uint8_t*)src.raw_data(), src.byte_size(), embed_ptr); - d.input_embeds_coords.emplace_back(size, p.begin() + (o.begin() - s.begin())); - } - } + if (offset + size > sum) { + /// TODO: reject for src range OOB + return Request::kInvalid; } + c.input_embeds_offsets.push_back(range.begin()); + c.input_embeds.push_back(embeds.slice(offset, size)); // reference into `embeds` + offset += size; + last = range.end(); } } - void Prepare(int phase, TensorMap& env) - { - auto& d = data_.at(phase); - auto& b = *env.at("batch").data()[0]; - auto& copy = *env.at("copy").data()[0]; - - // last output token + draft tokens - const Buffer_ autoreg_ids = env.at("autoreg_ids").buffer(); - - // core::CopyT copy{}; + return 0; +} - if (auto g = copy.group()) { - for (int i = 0; i < b.bsz; ++i) { - if (auto pos = d.autoreg_ids_pos[i]; pos >= 0) { - TM_CHECK_LT(b.perm[i], b.bs0); - copy(autoreg_ids.data() + b.perm[i], 1, &d.input_ids[pos]); - } - } +void InputProcessor::Impl::Add(int phase, TensorMap& env) +{ + const Buffer_ rc = env.at("requests").buffer(); + for (int i = 0; i < rc.size(); ++i) { + auto& c = *TM_CHECK_NOTNULL(rc[i]); + if (c.status == 0) { + c.status = Add(c); } - - env.produce("input_ids", d.input_ids.slice(0, d.input_token_num)); - env.produce("q_offsets", d.input_ids_offsets.slice(0, b.bsz + 1)); - env.produce("selected_token_pos", d.selected_token_pos.slice(0, b.bsz)); } +} - void PatchInputEmbedding(int phase, Tensor& embeds, BatchCopy& copy) - { - auto& d = data_.at(phase); - const auto byte_stride = byte_size(embeds.dtype(), embeds.stride(0)); - int offset = 0; - for (const auto& [size, pos] : d.input_embeds_coords) { - auto src = d.input_embeds_buf.slice(offset, size); - copy((uint8_t*)src.raw_data(), src.byte_size(), (uint8_t*)embeds.raw_data() + byte_stride * pos); - offset += size; - } +void InputProcessor::Impl::Setup(int phase, TensorMap& env) +{ + if (speculative_engine_) { + return SetupSpeculative(phase, env); } + return SetupOrdinary(phase, env); +} - void PatchMultimodalEmbedding(Tensor& embeds, BatchCopy& copy, const MultiModalEmbeddingData& multimodal) - { - TM_CHECK_EQ(multimodal.input_embeds_coords.size(), multimodal.image_embeds_coords.size()); - const int num_embeddings = multimodal.image_embeds_coords.size(); - for (int i = 0; i < num_embeddings; ++i) { - const auto& [sz0, image_offset] = multimodal.image_embeds_coords[i]; - const auto& [sz1, input_offset] = multimodal.input_embeds_coords[i]; - TM_CHECK_EQ(sz0, sz1); - copy(multimodal.data.slice(image_offset, sz0).buffer(), - sz0 * embeds.shape(1), - embeds.slice(input_offset, sz1).buffer()); +void InputProcessor::Impl::SetupOrdinary(int phase, TensorMap& env) +{ + auto& d = data_.at(phase); + auto& b = *env.at("batch").data()[0]; + auto& copy = *env.at("copy").data()[0]; + + Buffer_ rc = env.at("requests").buffer(); + + input_ids_offsets_buf_[0] = 0; + for (int i = 0; i < rc.size(); ++i) { + const Sequence& c = *rc[i]; + const SubmittedRow& row = *c.submitted; + const int q_begin = input_ids_offsets_buf_[i]; + const int q_len = row.input_len; + + input_ids_offsets_buf_[i + 1] = q_begin; + if (TM_UNLIKELY(!row.autoregres)) { + const int* src = c.token_ids + row.history_len + c.inflight_input_len; + std::copy_n(src, q_len, input_ids_buf_.data() + q_begin); + d.autoreg_ids_pos[i] = -1; + input_ids_offsets_buf_[i + 1] += q_len; } + else { + d.autoreg_ids_pos[i] = q_begin; + input_ids_offsets_buf_[i + 1] += 1; + } + decode_token_pos_buf_[i] = input_ids_offsets_buf_[i + 1] - 1; } - void PatchEmbedding(int phase, Tensor& embeds, BatchCopy& copy, TensorMap& env) - { - PatchInputEmbedding(phase, embeds, copy); + d.selected_token_count = rc.size(); - if (env.try_("multimodal")) { - const auto& multimodal = *env.at("multimodal").data()[0]; - PatchMultimodalEmbedding(embeds, copy, multimodal); - } + copy(input_ids_buf_, input_ids_offsets_buf_[b.bsz], d.input_ids); + copy(decode_token_pos_buf_, d.selected_token_count, d.selected_token_pos); + copy(input_ids_offsets_buf_, b.bsz + 1, d.input_ids_offsets); + + d.input_token_num = input_ids_offsets_buf_[b.bsz]; + + env.produce("token_num", Buffer{&d.input_token_num, 1, kCPU}); + + StageEmbeddingPatches(phase, rc, false); +} + +void InputProcessor::Impl::Prepare(int phase, TensorMap& env) +{ + if (speculative_engine_) { + return PrepareSpeculative(phase, env); } + return PrepareOrdinary(phase, env); +} -private: - struct Data { - Buffer_ input_ids; - Buffer_ input_ids_offsets; - int input_token_num; +void InputProcessor::Impl::PrepareOrdinary(int phase, TensorMap& env) +{ + auto& d = data_.at(phase); + auto& b = *env.at("batch").data()[0]; + auto& copy = *env.at("copy").data()[0]; - Buffer_ selected_token_pos; + const Buffer_ autoreg_ids = env.at("autoreg_ids").buffer(); - Buffer_ autoreg_ids_pos; + if (auto g = copy.group()) { + for (int i = 0; i < b.bsz; ++i) { + if (auto pos = d.autoreg_ids_pos[i]; pos >= 0) { + TM_CHECK_LT(b.perm[i], b.bs0); + copy(autoreg_ids.data() + b.perm[i], 1, &d.input_ids[pos]); + } + } + } - Tensor input_embeds_buf; - vector> input_embeds_coords; // (size, pos) - }; + env.produce("input_ids", d.input_ids.slice(0, d.input_token_num)); + env.produce("q_offsets", d.input_ids_offsets.slice(0, b.bsz + 1)); + env.produce("selected_token_pos", d.selected_token_pos.slice(0, d.selected_token_count)); +} -private: - const int max_batch_size_; - const int max_forward_token_num_; +void InputProcessor::Impl::StageEmbeddingPatches(int phase, + const Buffer_& rc, + bool stage_successor) +{ + auto& d = data_.at(phase); + + d.target_input_patches.clear(); + d.successor_input_patches.clear(); + auto embed_ptr = (uint8_t*)d.input_embeds_buf.raw_data(); + int staged_row = 0; + for (int k = 0; k < rc.size(); ++k) { + auto& c = *rc[k]; + const SubmittedRow& row = *c.submitted; + if (!row.autoregres) { + const auto& embeds = c.input_embeds; + const auto& offsets = c.input_embeds_offsets; + const int begin = row.history_len + c.inflight_input_len; + const int end = begin + row.input_len; + const Interval target{begin, end}; + const Interval successor{begin + 1, std::min(end + 1, c.seq_len)}; + const Interval staged{begin, std::min(end + 1, c.seq_len)}; + const int packed_begin = input_ids_offsets_buf_[k]; + for (int i = (int)offsets.size() - 1; i >= 0; --i) { + Interval r{offsets[i], Interval::Size{(int)embeds[i].shape(0)}}; + auto staged_overlap = r & staged; + if (auto size = (int)staged_overlap.size()) { + auto src = embeds[i].slice(staged_overlap.begin() - r.begin(), size); + embed_ptr = std::copy_n((const uint8_t*)src.raw_data(), src.byte_size(), embed_ptr); + + if (auto o = r & target; !o.empty()) { + d.target_input_patches.push_back( + {static_cast(o.size()), + staged_row + o.begin() - staged_overlap.begin(), + packed_begin + o.begin() - target.begin()}); + } + if (stage_successor) { + if (auto o = r & successor; !o.empty()) { + d.successor_input_patches.push_back( + {static_cast(o.size()), + staged_row + o.begin() - staged_overlap.begin(), + packed_begin + o.begin() - successor.begin()}); + } + } + staged_row += size; + } + } + } + } +} - vector data_; +void InputProcessor::Impl::ApplyPatches(const Tensor& source, + const std::vector& patches, + Tensor& destination, + BatchCopy& copy) +{ + for (const EmbeddingPatch& patch : patches) { + copy(source.slice(patch.source_row, patch.row_count).buffer(), + patch.row_count * destination.shape(1), + destination.slice(patch.destination_row, patch.row_count).buffer()); + } +} - Buffer_ input_ids_buf_; - Buffer_ input_ids_offsets_buf_; +void InputProcessor::Impl::PatchEmbedding(int phase, Tensor& embeds, BatchCopy& copy, TensorMap& env) +{ + auto& data = data_.at(phase); + ApplyPatches(data.input_embeds_buf, data.target_input_patches, embeds, copy); + if (env.try_("multimodal")) { + const auto& multimodal = *env.at("multimodal").data()[0]; + ApplyPatches(multimodal.data, multimodal.target_patches, embeds, copy); + } +} - Buffer_ decode_token_pos_buf_; -}; +void InputProcessor::Impl::PatchSuccessorEmbedding(int phase, Tensor& embeds, BatchCopy& copy, TensorMap& env) +{ + auto& data = data_.at(phase); + ApplyPatches(data.input_embeds_buf, data.successor_input_patches, embeds, copy); + if (env.try_("multimodal")) { + const auto& multimodal = *env.at("multimodal").data()[0]; + ApplyPatches(multimodal.data, multimodal.successor_patches, embeds, copy); + } +} InputProcessor::~InputProcessor() = default; -InputProcessor::InputProcessor(const EngineParam& engine, int hidden_units, DataType data_type, int phases): - impl_{std::make_unique(engine, hidden_units, data_type, phases)} +InputProcessor::InputProcessor(const EngineParam& engine, + int hidden_units, + DataType data_type, + int phases, + bool speculative, + int max_verification_positions, + bool successor_embeddings): + impl_{std::make_unique( + engine, hidden_units, data_type, phases, speculative, max_verification_positions, successor_embeddings)} { } @@ -278,9 +313,19 @@ void InputProcessor::Run(BatchOp op, int phase, TensorMap& env) } } +void InputProcessor::BuildTargetInputs(int phase, TensorMap& env) +{ + impl_->BuildTargetInputs(phase, env); +} + void InputProcessor::PatchEmbedding(int phase, Tensor& embeds, BatchCopy& copy, TensorMap& env) { impl_->PatchEmbedding(phase, embeds, copy, env); } +void InputProcessor::PatchSuccessorEmbedding(int phase, Tensor& embeds, BatchCopy& copy, TensorMap& env) +{ + impl_->PatchSuccessorEmbedding(phase, embeds, copy, env); +} + } // namespace turbomind diff --git a/src/turbomind/models/input_processor.h b/src/turbomind/models/input_processor.h index cee930fc72..07ac93b43d 100644 --- a/src/turbomind/models/input_processor.h +++ b/src/turbomind/models/input_processor.h @@ -9,12 +9,25 @@ class InputProcessor { public: ~InputProcessor(); - InputProcessor(const EngineParam& engine, int hidden_units, DataType data_type, int phases); + InputProcessor(const EngineParam& engine, + int hidden_units, + DataType data_type, + int phases, + bool speculative, + int max_verification_positions, + bool successor_embeddings); void Run(BatchOp op, int phase, TensorMap& env); + // Composed mode only: the forward-time target-pass staging step. Gathers + // the target's input ids from the request token rows and stages the + // per-request key lengths for the executor's offsets prefix-sum. + void BuildTargetInputs(int phase, TensorMap& env); + void PatchEmbedding(int phase, Tensor& embeds, BatchCopy& copy, TensorMap& env); + void PatchSuccessorEmbedding(int phase, Tensor& embeds, BatchCopy& copy, TensorMap& env); + private: struct Impl; std::unique_ptr impl_; diff --git a/src/turbomind/models/input_processor_impl.h b/src/turbomind/models/input_processor_impl.h new file mode 100644 index 0000000000..e4bc5d4229 --- /dev/null +++ b/src/turbomind/models/input_processor_impl.h @@ -0,0 +1,88 @@ +// Copyright (c) OpenMMLab. All rights reserved. +#pragma once + +#include + +#include "src/turbomind/core/core.h" +#include "src/turbomind/engine/batch.h" +#include "src/turbomind/engine/request.h" +#include "src/turbomind/models/input_processor.h" +#include "src/turbomind/models/vision_model.h" + +namespace turbomind { + +struct InputProcessor::Impl { + + Impl(const EngineParam& engine, + int hidden_units, + DataType data_type, + int phases, + bool speculative, + int max_verification_positions, + bool successor_embeddings); + + int Add(Sequence& c); + void Add(int phase, TensorMap& env); + + // Shared lifecycle entry: dispatches to the mode implementation. + void Setup(int phase, TensorMap& env); + void Prepare(int phase, TensorMap& env); + + // Ordinary mode (input_processor.cc): generating rows trim to one token, + // autoregressive ids patch the packed row, and selected positions are + // row-major. + void SetupOrdinary(int phase, TensorMap& env); + void PrepareOrdinary(int phase, TensorMap& env); + + // Composed mode (input_processor_speculative.cc): all query rows pack + // verbatim, selected positions carry the verification geometry, and + // target ids are gathered from the request's token row. + void SetupSpeculative(int phase, TensorMap& env); + void PrepareSpeculative(int phase, TensorMap& env); + void BuildTargetInputs(int phase, TensorMap& env); + + // Stages input-embedding patches for the submitted rows; successor + // staging is requested only by the composed mode. + void StageEmbeddingPatches(int phase, const Buffer_& rc, bool stage_successor); + + static void ApplyPatches(const Tensor& source, + const std::vector& patches, + Tensor& destination, + BatchCopy& copy); + + void PatchEmbedding(int phase, Tensor& embeds, BatchCopy& copy, TensorMap& env); + void PatchSuccessorEmbedding(int phase, Tensor& embeds, BatchCopy& copy, TensorMap& env); + + struct Data { + Buffer_ input_ids; + Buffer_ input_ids_offsets; + int input_token_num; + + Buffer_ selected_token_pos; + int selected_token_count; + Buffer_ target_ids_from_row; + + Buffer_ target_key_lengths; + + Buffer_ autoreg_ids_pos; + + Tensor input_embeds_buf; + std::vector target_input_patches; + std::vector successor_input_patches; + }; + + const int max_batch_size_; + const int max_forward_token_num_; + const bool speculative_engine_; + const bool successor_embeddings_; + + std::vector data_; + + Buffer_ input_ids_buf_; + Buffer_ input_ids_offsets_buf_; + + Buffer_ decode_token_pos_buf_; + Buffer_ target_ids_from_row_buf_; +}; + +} // namespace turbomind diff --git a/src/turbomind/models/input_processor_speculative.cc b/src/turbomind/models/input_processor_speculative.cc new file mode 100644 index 0000000000..1ed76cee13 --- /dev/null +++ b/src/turbomind/models/input_processor_speculative.cc @@ -0,0 +1,104 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#include "src/turbomind/core/check.h" +#include "src/turbomind/core/context.h" +#include "src/turbomind/core/core.h" + +#include "src/turbomind/engine/request.h" + +#include "src/turbomind/kernels/speculative_sequence_kernels.h" + +#include "src/turbomind/models/input_processor_impl.h" + +namespace turbomind { + +void InputProcessor::Impl::SetupSpeculative(int phase, TensorMap& env) +{ + auto& d = data_.at(phase); + auto& b = *env.at("batch").data()[0]; + auto& copy = *env.at("copy").data()[0]; + + Buffer_ rc = env.at("requests").buffer(); + + input_ids_offsets_buf_[0] = 0; + int generation_count = 0; + for (int i = 0; i < rc.size(); ++i) { + const Sequence& c = *rc[i]; + const SubmittedRow& row = *c.submitted; + const int q_begin = input_ids_offsets_buf_[i]; + const int q_len = row.input_len; + + input_ids_offsets_buf_[i + 1] = q_begin + q_len; + target_ids_from_row_buf_[i] = row.autoregres; + if (!row.autoregres) { + const int* src = c.token_ids + row.history_len + c.inflight_input_len; + std::copy_n(src, q_len, input_ids_buf_.data() + q_begin); + } + if (row.generating) { + ++generation_count; + } + } + + const int position_count = env.at("verification_positions").data()[0]; + + int g = 0; + for (int i = 0; i < rc.size(); ++i) { + const Sequence& c = *rc[i]; + const SubmittedRow& submitted = *c.submitted; + if (!submitted.generating) { + continue; + } + + const int q_end = input_ids_offsets_buf_[i + 1]; + + for (int position = 0; position < position_count; ++position) { + // Wide rows carry width == input_len; positions beyond a row's width clamp to its last token. + decode_token_pos_buf_[position * generation_count + g] = + q_end - 1 - std::max(0, submitted.verification_positions - 1 - position); + } + ++g; + } + + d.selected_token_count = position_count * generation_count; + + copy(input_ids_buf_, input_ids_offsets_buf_[b.bsz], d.input_ids); + copy(decode_token_pos_buf_, d.selected_token_count, d.selected_token_pos); + copy(input_ids_offsets_buf_, b.bsz + 1, d.input_ids_offsets); + copy(target_ids_from_row_buf_, b.bsz, d.target_ids_from_row); + + d.input_token_num = input_ids_offsets_buf_[b.bsz]; + + env.produce("token_num", Buffer{&d.input_token_num, 1, kCPU}); + + StageEmbeddingPatches(phase, rc, successor_embeddings_); +} + +void InputProcessor::Impl::PrepareSpeculative(int phase, TensorMap& env) +{ + auto& d = data_.at(phase); + auto& b = *env.at("batch").data()[0]; + + env.produce("input_ids", d.input_ids.slice(0, d.input_token_num)); + env.produce("q_offsets", d.input_ids_offsets.slice(0, b.bsz + 1)); + env.produce("selected_token_pos", d.selected_token_pos.slice(0, d.selected_token_count)); +} + +void InputProcessor::Impl::BuildTargetInputs(int phase, TensorMap& env) +{ + auto& d = data_.at(phase); + auto& b = *env.at("batch").data()[0]; + + invokeBuildTargetInputs(env.at("input_ids").data(), + d.target_key_lengths.data(), + reinterpret_cast(env.at("request_token_ids_ptrs").data()), + env.at("q_offsets").data(), + env.at("sequence_length").data(), + d.target_ids_from_row.data(), + env.at("finished").data(), + b.bsz, + core::Context::stream().handle()); + + env.produce("target_key_lengths", d.target_key_lengths.slice(0, b.bsz)); +} + +} // namespace turbomind diff --git a/src/turbomind/models/internvit/internvit.cc b/src/turbomind/models/internvit/internvit.cc index b0f632ac2e..82e30d4e72 100644 --- a/src/turbomind/models/internvit/internvit.cc +++ b/src/turbomind/models/internvit/internvit.cc @@ -58,14 +58,15 @@ struct InternVit::Impl { const int tp_group_; const int tp_size_; const DataType engine_data_type_; + const bool successor_embeddings_; Buffer_ attn_cu_seqlens_buf_; struct Data { Tensor batch_input; int batch_size{}; - std::vector> image_embeds_coords; - std::vector> input_embeds_coords; + std::vector target_patches; + std::vector successor_patches; Tensor_ attn_cu_seqlens; Tensor_ attn_finished; int token_num{}; @@ -76,14 +77,18 @@ struct InternVit::Impl { batch_size = 0; token_num = 0; seq_len = 0; - image_embeds_coords.clear(); - input_embeds_coords.clear(); + target_patches.clear(); + successor_patches.clear(); } }; std::vector data_; - Impl(const EngineParam& engine, const Context& ctx, const InternVitWeight& weights, int phases): + Impl(const EngineParam& engine, + const Context& ctx, + const InternVitWeight& weights, + int phases, + bool successor_embeddings): weights_{weights}, config_{weights.config()}, h_tp_group{ctx.comm.h_comm}, @@ -91,7 +96,8 @@ struct InternVit::Impl { d_comm_{ctx.comm.d_comm}, tp_group_{ctx.comm.d_tp_group}, tp_size_{ctx.comm.h_tp_group ? ctx.comm.h_tp_group->n_ranks() : 1}, - engine_data_type_{engine.data_type} + engine_data_type_{engine.data_type}, + successor_embeddings_{successor_embeddings} { const auto& cfg = weights.config(); for (int i = 0; i < phases; ++i) { @@ -239,29 +245,42 @@ struct InternVit::Impl { Buffer_ rc = env.at("requests").buffer(); for (int i = 0; i < rc.size(); ++i) { - const auto& s = *rc[i]; + const Sequence& s = *rc[i]; + const SubmittedRow& submitted = *s.submitted; - if ((not s.autoregres) && (not s.multimodal_inputs.empty())) { + if ((not submitted.autoregres) && (not s.multimodal_inputs.empty())) { ++mm_prefill_seqs; images_total += (int)s.multimodal_inputs.size(); - Interval text{s.history_len + s.inflight_input_len, Interval::Size{s.input_len}}; + const int begin = submitted.history_len + s.inflight_input_len; + const int end = begin + submitted.input_len; + const Interval target{begin, end}; + const Interval successor{begin + 1, std::min(end + 1, s.seq_len)}; for (const auto& mm : s.multimodal_inputs) { - auto o = mm->interval & text; - if (auto size = (int)o.size()) { + const Interval target_overlap = mm->interval & target; + const Interval successor_overlap = successor_embeddings_ ? mm->interval & successor : Interval{}; + if (!target_overlap.empty() || !successor_overlap.empty()) { pixel_values.push_back(mm->data); d.batch_size += mm->data.shape(0); - const int text_offset = input_ids_offsets + o.begin() - text.begin(); - const int image_offset = image_embeds_offsets + o.begin() - mm->interval.begin(); - d.input_embeds_coords.emplace_back(size, text_offset); - d.image_embeds_coords.emplace_back(size, image_offset); + if (!target_overlap.empty()) { + d.target_patches.push_back( + {static_cast(target_overlap.size()), + image_embeds_offsets + target_overlap.begin() - mm->interval.begin(), + input_ids_offsets + target_overlap.begin() - target.begin()}); + } + if (!successor_overlap.empty()) { + d.successor_patches.push_back( + {static_cast(successor_overlap.size()), + image_embeds_offsets + successor_overlap.begin() - mm->interval.begin(), + input_ids_offsets + successor_overlap.begin() - successor.begin()}); + } image_embeds_offsets += (int)mm->interval.size(); } } } - input_ids_offsets += s.autoregres ? 1 : s.input_len; + input_ids_offsets += submitted.input_len; } // Prefix-cache observability: on a fully-cached image, the window filter @@ -613,13 +632,16 @@ struct InternVit::Impl { EnsureFloatDtype(image_embeds, engine_data_type_); - args.produce("multimodal", - MultiModalEmbeddingData{image_embeds, d.image_embeds_coords, d.input_embeds_coords}.buf()); + args.produce("multimodal", MultiModalEmbeddingData{image_embeds, d.target_patches, d.successor_patches}.buf()); } }; -InternVit::InternVit(const EngineParam& engine, const Context& ctx, const InternVitWeight& weights, int phases): - impl_{std::make_unique(engine, ctx, weights, phases)} +InternVit::InternVit(const EngineParam& engine, + const Context& ctx, + const InternVitWeight& weights, + int phases, + bool successor_embeddings): + impl_{std::make_unique(engine, ctx, weights, phases, successor_embeddings)} { } diff --git a/src/turbomind/models/internvit/internvit.h b/src/turbomind/models/internvit/internvit.h index 5d5ba46b9e..c31f5e2f73 100644 --- a/src/turbomind/models/internvit/internvit.h +++ b/src/turbomind/models/internvit/internvit.h @@ -12,7 +12,11 @@ class InternVitWeight; class InternVit: public VisionModel { public: - InternVit(const EngineParam& engine, const Context& ctx, const InternVitWeight& weights, int phases); + InternVit(const EngineParam& engine, + const Context& ctx, + const InternVitWeight& weights, + int phases, + bool successor_embeddings); ~InternVit() override; diff --git a/src/turbomind/models/language_model.cc b/src/turbomind/models/language_model.cc index 0555f1c45d..ef395b0efc 100644 --- a/src/turbomind/models/language_model.cc +++ b/src/turbomind/models/language_model.cc @@ -1,40 +1,30 @@ #include "src/turbomind/models/language_model.h" +#include #include #include #include "src/turbomind/comm/device_comm.h" +#include "src/turbomind/comm/host_comm.h" #include "src/turbomind/core/allocator.h" #include "src/turbomind/core/check.h" #include "src/turbomind/core/context.h" #include "src/turbomind/core/copy.h" -#include "src/turbomind/core/interval.h" #include "src/turbomind/core/scope.h" -#include "src/turbomind/core/state.h" -#include "src/turbomind/engine/batch.h" #include "src/turbomind/engine/cache_registry.h" -#include "src/turbomind/engine/request.h" -#include "src/turbomind/generation/generation.h" #include "src/turbomind/kernels/gpt_kernels.h" -#include "src/turbomind/models/input_processor.h" #include "src/turbomind/models/llama/llama_kernels.h" #include "src/turbomind/models/llama/llama_params.h" -#include "src/turbomind/models/llama/llama_utils.h" #include "src/turbomind/models/llama/unified_decoder.h" #include "src/turbomind/models/model_weight.h" -#include "src/turbomind/models/output_processor.h" -#include "src/turbomind/utils/anomaly_handler.h" #include "src/turbomind/utils/cuda_utils.h" +#include "src/turbomind/utils/nvtx_utils.h" // #include "dbg.h" namespace turbomind { -using std::vector; -using std::unique_ptr; -using std::shared_ptr; - struct LanguageModel::Impl { const Communicators& comm_; const ModelWeight& weights_; @@ -44,540 +34,363 @@ struct LanguageModel::Impl { const int tp_rank_; const bool use_ag2d_; - const int attn_dp_size_; - const int attn_dp_rank_; - const int max_batch_size_; - - const bool debug_; - - Buffer_ false_; - - // mutable state - State finished_; - State sequence_length_; // length of known tokens - // immutable state - Buffer_ autoreg_ids_; - // Buffer_ autoreg_ids_offsets_; - - // Symmetric buffer for holding global hidden states or logits - Buffer_ symm_buf_; - - // Global (all attention DP ranks) per-token validity mask, built at Forward time; - // consumed by the attention layers (their DP-local slice) and, eventually, the MoE router. - Buffer_ token_mask_; - - // Symmetric gather buffer for the per-rank `[q_offsets | finished]` metadata blocks - // ([attn_dp_size, meta_bytes], 16B-aligned rows for the in-place AllGather). - // Only allocated when attn_dp > 1. - Tensor_ symm_token_meta_; - // Max chunk size for compute / output full logits int max_logits_len_ = 0; - Buffer_ sequence_length_buf_; - Buffer_ readonly_block_num_buf_; // {max_batch_size}, kCPUpinned - Buffer_ finished_buf_; - - struct Data { - Buffer_ sequence_length; - Buffer_ readonly_block_num; - Buffer_ finished; - - Buffer_ autoregres; - Buffer_ generating; - - int n_generating; - }; - - vector data_; - - std::optional input_processor_; std::unique_ptr unified_decoder_; - std::optional output_processor_; - std::unique_ptr generation_; // token generator void Run(BatchOp op, int phase, TensorMap& env) { - switch (op) { - case BatchOp::kSetup: - return Setup(phase, env); - case BatchOp::kPrepare: - return Prepare(phase, env); - case BatchOp::kForward: - return Forward(phase, env); - case BatchOp::kUnprep: - return Unprep(phase, env); - case BatchOp::kFetch: - return Fetch(phase, env); - default: - input_processor_->Run(op, phase, env); - unified_decoder_->Run(op, phase, env); - generation_->Run(op, phase, env); - output_processor_->Run(op, phase, env); - } + unified_decoder_->Run(op, phase, env); } - Impl( - CacheRegistry& registry, const EngineParam& engine, const Context& ctx, const ModelWeight& weights, int phases); - - Tensor LookupEmbedding(const Buffer_& input_ids, Buffer symm_buf); - Tensor PostEmbedding(const Tensor& features, Buffer symm_buf); - - // Build the global per-token validity mask for this pass (see `token_mask_`). - void BuildTokenMask(const bool* finished, const int* q_offsets, const BatchData& b); + Impl(CacheRegistry& registry, + const EngineParam& engine, + const Context& ctx, + const ModelWeight& weights, + int phases); + + Tensor LookupEmbedding(const Buffer_& input_ids, + const Tensor& embedding_table, + Buffer model_tp_gather_buffer, + Tensor embeddings); + Tensor PostEmbedding(const Tensor& features, + const LinearWeight& output_weight, + Buffer model_tp_gather_buffer, + Tensor logits); + + bool has_embedding() const; + bool has_head() const; + Tensor Embed(const Buffer_& input_ids, Tensor out, const TensorMap& env); + LanguageModel::DecoderOutputs RunDecoder(int phase, const LanguageModel::DecoderInputs& in, TensorMap& env); + Tensor Logits(const Tensor& hidden, Tensor out, const TensorMap& env); + const ModelWeight& weights() const; + int max_logits_len(const TensorMap& env) const; + bool logits_use_workspace() const; + void CommitAcceptedState(int phase, const Buffer_& accept_len); + size_t SpeculativeStateJournalBytes(int request_count, int verification_positions) const; - void Setup(int phase, TensorMap& env); - void Prepare(int phase, TensorMap& env); - void Forward(int phase, TensorMap& env); - void Unprep(int phase, TensorMap& env); - void Fetch(int phase, TensorMap& env); }; -LanguageModel::Impl::Impl( - CacheRegistry& registry, const EngineParam& engine, const Context& ctx, const ModelWeight& weights, int phases): +LanguageModel::Impl::Impl(CacheRegistry& registry, + const EngineParam& engine, + const Context& ctx, + const ModelWeight& weights, + int phases): comm_{ctx.comm}, weights_{weights}, linear_{*ctx.linear}, tp_size_{comm_.h_tp_group->n_ranks()}, tp_rank_{comm_.h_tp_group->rank()}, - use_ag2d_{comm_.d_comm && comm_.d_comm->Query(comm::kHasAllGather2D)}, - attn_dp_size_{engine.attn_dp_size}, - attn_dp_rank_{engine.attn_dp_rank}, - max_batch_size_{engine.max_batch_size}, - debug_{isDebug()} + use_ag2d_{comm_.d_comm && comm_.d_comm->Query(comm::kHasAllGather2D)} { - - false_ = {engine.max_batch_size, kDEVICE}; - Clear(false_); - - finished_buf_ = {engine.max_batch_size, kCPUpinned}; - finished_ = {{engine.max_batch_size}, kBool, kDEVICE}; - - autoreg_ids_ = {engine.max_batch_size, kDEVICE}; - // autoreg_ids_offsets_ = {engine.max_batch_size + 1, kCPU}; - // std::fill_n(autoreg_ids_offsets_.data(), autoreg_ids_offsets_.size(), 0); - - sequence_length_buf_ = {engine.max_batch_size, kCPUpinned}; - readonly_block_num_buf_ = {engine.max_batch_size, kCPUpinned}; - sequence_length_ = {{engine.max_batch_size}, kInt, kDEVICE}; - for (int i = 0; i < phases; ++i) { - auto& d = data_.emplace_back(); - d.sequence_length = empty_like(sequence_length_buf_, kDEVICE); - d.readonly_block_num = empty_like(readonly_block_num_buf_, kDEVICE); - d.finished = empty_like(finished_buf_, kDEVICE); - d.autoregres = {engine.max_batch_size, kCPU}; - d.generating = {engine.max_batch_size, kCPU}; - } - - input_processor_.emplace(engine, weights_.hidden_units, weights_.data_type, phases); + const int max_logits_rows = + engine.max_batch_size * (engine.spec_method.empty() ? 1 : engine.spec_num_draft_tokens + 1); unified_decoder_ = std::make_unique(registry, engine, ctx, phases, weights_); - const int vocab_size = weights_.output->output_dim * tp_size_; - - generation_ = std::make_unique( - kFloat32, engine.max_batch_size, engine.session_len, weights_.vocab_size, vocab_size, comm_.h_tp_group, phases); - - const ssize_t max_fwd_tokens = engine.max_forward_token_num; - - if (ctx.comm.d_comm) { - auto symm_alloc = GetSymmAllocator(ctx.comm.d_comm); - // Native comm fuses allreduce & rmsnorm in token granularity - TM_CHECK(engine.max_forward_token_num % tp_size_ == 0); - - ssize_t bytes{}; - bytes = std::max(bytes, - byte_size(weights_.data_type, max_fwd_tokens * engine.attn_dp_size * weights_.hidden_units)); - bytes = std::max(bytes, byte_size(weights_.data_type, engine.max_batch_size * vocab_size)); - - symm_buf_ = {bytes, symm_alloc}; - // Compute max logits length based on symm buffer size - max_logits_len_ = symm_buf_.view(weights_.data_type).size() / vocab_size; - - if (attn_dp_size_ > 1) { - const int q_bytes = (max_batch_size_ + 1) * (int)sizeof(int); - const int meta_bytes = (q_bytes + max_batch_size_ + 15) / 16 * 16; - symm_token_meta_ = {{attn_dp_size_, meta_bytes}, symm_alloc}; - } - } - else { - max_logits_len_ = std::max(max_fwd_tokens * weights_.hidden_units / vocab_size, engine.max_batch_size); + if (has_head()) { + TM_CHECK_GT(weights_.vocab_size_padded, 0) << "a model without a head has no logits length"; + max_logits_len_ = std::max( + core::ssize_t(engine.max_forward_token_num) * weights_.hidden_units / weights_.vocab_size_padded, + max_logits_rows); } - - token_mask_ = {max_fwd_tokens * attn_dp_size_, kDEVICE}; - - output_processor_.emplace(weights_.vocab_size, max_logits_len_, tp_rank_, phases, [this](const Tensor& hstate) { - return PostEmbedding(hstate, symm_buf_); - }); } -Tensor LanguageModel::Impl::LookupEmbedding(const Buffer_& input_ids, Buffer symm_buf) +Tensor LanguageModel::Impl::LookupEmbedding(const Buffer_& input_ids, + const Tensor& embedding_table, + Buffer model_tp_gather_buffer, + Tensor embeddings) { TM_FUNCTION_SCOPE(); - const auto st = core::Context::stream().handle(); - - const int hidden_units = weights_.hidden_units; - const auto& embedding_table = weights_.tok_embeddings; - TM_CHECK_EQ(embedding_table.shape(1) * tp_size_, hidden_units); + const int token_count = input_ids.size(); + const core::Device expected_device = core::Context::device_alloc()->device(); + const bool embeddings_supplied = embeddings.ndim() != 0; + const int local_hidden_size = static_cast(embedding_table.shape(1)); + const int hidden_size = local_hidden_size * tp_size_; + const DataType data_type = embedding_table.dtype(); - const int token_num = input_ids.size(); - - Tensor input_embeds{{token_num, hidden_units}, weights_.data_type, kDEVICE}; + if (token_count == 0) { + if (embeddings_supplied) { + return embeddings; + } + return Tensor{static_cast(nullptr), Layout{{0, hidden_size}}, data_type, expected_device}; + } - if (token_num == 0) { - return input_embeds; + if (!embeddings_supplied) { + embeddings = Tensor{{token_count, hidden_size}, data_type, expected_device}; } + const cudaStream_t stream = core::Context::stream().handle(); + if (tp_size_ == 1) { - invokeEmbeddingLookup(input_embeds, input_ids, embedding_table, st); - TM_CUDA_CHECK(cudaGetLastError()); + invokeEmbeddingLookup(embeddings, input_ids, embedding_table, stream); } else if (use_ag2d_) { - const auto local_hidden_units = embedding_table.shape(1); + Tensor gathered{model_tp_gather_buffer.view(data_type), {token_count, tp_size_, local_hidden_size}}; + Tensor local = gathered.slice({0, tp_rank_, 0}, {token_count, 1, local_hidden_size}).squeeze(1); + Tensor flat = gathered.view({token_count, hidden_size}); - Tensor temp{symm_buf.view(weights_.data_type), {token_num, tp_size_, local_hidden_units}}; - Tensor local{temp.slice({0, tp_rank_, 0}, {-1, 1, -1}).squeeze(1)}; - - invokeEmbeddingLookup(local, input_ids, embedding_table, st); - TM_CUDA_CHECK(cudaGetLastError()); + const bool exact_alias = embeddings.raw_data() == flat.raw_data(); + invokeEmbeddingLookup(local, input_ids, embedding_table, stream); comm_.d_comm->AllGather2D(local.raw_data(), - temp.raw_data(), - hidden_units, - local_hidden_units, - local_hidden_units, - token_num, - local.dtype(), + gathered.raw_data(), + hidden_size, + local_hidden_size, + local_hidden_size, + token_count, + gathered.dtype(), {true, true}, comm_.d_tp_group, - st); - TM_CUDA_CHECK(cudaGetLastError()); + stream); - Copy(temp.buffer(), input_embeds.buffer()); + if (!exact_alias) { + Copy(flat, embeddings); + } } else { - const auto local_hidden_units = embedding_table.shape(1); - - Tensor temp{symm_buf.view(weights_.data_type), {tp_size_, token_num, local_hidden_units}}; - Tensor local{temp.slice(tp_rank_).squeeze(0)}; - - invokeEmbeddingLookup(local, input_ids, embedding_table, st); - TM_CUDA_CHECK(cudaGetLastError()); + Tensor gathered{model_tp_gather_buffer.view(data_type), {tp_size_, token_count, local_hidden_size}}; + Tensor local = gathered.slice({tp_rank_, 0, 0}, {1, token_count, local_hidden_size}).squeeze(0); + invokeEmbeddingLookup(local, input_ids, embedding_table, stream); comm_.d_comm->AllGather( - local.raw_data(), temp.raw_data(), local.size(), weights_.data_type, comm_.d_tp_group, st); - TM_CUDA_CHECK(cudaGetLastError()); - - invokeInPlaceTranspose102((uint16_t*)input_embeds.raw_data(), - (uint16_t*)temp.raw_data(), + local.raw_data(), gathered.raw_data(), local.size(), data_type, comm_.d_tp_group, stream); + invokeInPlaceTranspose102(static_cast(embeddings.raw_data()), + static_cast(gathered.raw_data()), tp_size_, - token_num, - local_hidden_units, + token_count, + local_hidden_size, false, - st); - TM_CUDA_CHECK(cudaGetLastError()); + stream); } - return input_embeds; + return embeddings; } -Tensor LanguageModel::Impl::PostEmbedding(const Tensor& features, Buffer symm_buf) +Tensor LanguageModel::Impl::PostEmbedding(const Tensor& features, + const LinearWeight& output_weight, + Buffer model_tp_gather_buffer, + Tensor logits) { TM_FUNCTION_SCOPE(); NvtxScope scope("postDecodeEmbedding"); - const auto st = core::Context::stream().handle(); + const core::Device expected_device = core::Context::device_alloc()->device(); + const bool logits_supplied = logits.ndim() != 0; + + const int batch_size = static_cast(features.shape(0)); + const int local_vocab_size = output_weight.output_dim; + const int padded_vocab_size = local_vocab_size * tp_size_; + const DataType output_dtype = output_weight.output_dtype(); - const int bsz = features.shape(0); - const int local_vocab_size = weights_.output->output_dim; - const int vocab_size = local_vocab_size * tp_size_; + if (batch_size == 0) { + if (logits_supplied) { + return logits; + } + return Tensor{static_cast(nullptr), Layout{{0, padded_vocab_size}}, output_dtype, expected_device}; + } - if (bsz == 0) { - return Tensor{{0, vocab_size}, weights_.data_type, kDEVICE}; + if (!logits_supplied) { + if (tp_size_ > 1 && use_ag2d_) { + logits = Tensor{model_tp_gather_buffer.view(output_dtype), {batch_size, padded_vocab_size}}; + } + else { + logits = Tensor{{batch_size, padded_vocab_size}, output_dtype, expected_device}; + } } + const cudaStream_t stream = core::Context::stream().handle(); + if (tp_size_ == 1) { - Tensor logits{{bsz, vocab_size}, weights_.data_type, kDEVICE}; - TM_SCOPE_CALL(linear_.Forward(features, *weights_.output, logits)); - TM_DEBUG_TENSOR(logits, "logits", 1); - return logits; + TM_SCOPE_CALL(linear_.Forward(features, output_weight, logits)); } else if (use_ag2d_) { - Tensor logits{symm_buf.view(weights_.data_type), {bsz, tp_size_, local_vocab_size}}; - Tensor local = logits.slice({0, tp_rank_, 0}, {-1, 1, -1}); - TM_SCOPE_CALL(linear_.Forward(features, *weights_.output, local.squeeze(1))); + Tensor gathered = logits.view({batch_size, tp_size_, local_vocab_size}); + Tensor local = gathered.slice({0, tp_rank_, 0}, {batch_size, 1, local_vocab_size}).squeeze(1); + + TM_SCOPE_CALL(linear_.Forward(features, output_weight, local)); comm_.d_comm->AllGather2D(local.raw_data(), - logits.raw_data(), - vocab_size, + gathered.raw_data(), + padded_vocab_size, local_vocab_size, local_vocab_size, - bsz, - logits.dtype(), + batch_size, + gathered.dtype(), {true, true}, comm_.d_tp_group, - st); - TM_CUDA_CHECK(cudaGetLastError()); - return logits.view({bsz, -1}); + stream); } else { - Tensor logits{symm_buf.view(weights_.data_type), {tp_size_, bsz, local_vocab_size}}; - Tensor local = logits.slice({tp_rank_, 0, 0}, {1, -1, -1}); - TM_SCOPE_CALL(linear_.Forward(features, *weights_.output, local.squeeze(0))); - comm_.d_comm->AllGather(local.raw_data(), logits.raw_data(), local.size(), local.dtype(), comm_.d_tp_group, st); - TM_CUDA_CHECK(cudaGetLastError()); - Tensor out{{bsz, vocab_size}, features.dtype(), features.device()}; - invokeTransposeAxis01( - (uint16_t*)out.raw_data(), (uint16_t*)logits.raw_data(), tp_size_, bsz, local_vocab_size, st); - TM_CUDA_CHECK(cudaGetLastError()); - return out; + Tensor gathered{model_tp_gather_buffer.view(output_dtype), {tp_size_, batch_size, local_vocab_size}}; + Tensor local = gathered.slice({tp_rank_, 0, 0}, {1, batch_size, local_vocab_size}).squeeze(0); + + TM_SCOPE_CALL(linear_.Forward(features, output_weight, local)); + comm_.d_comm->AllGather( + local.raw_data(), gathered.raw_data(), local.size(), local.dtype(), comm_.d_tp_group, stream); + invokeTransposeAxis01(static_cast(logits.raw_data()), + static_cast(gathered.raw_data()), + tp_size_, + batch_size, + local_vocab_size, + stream); } + + return logits; } -void LanguageModel::Impl::Setup(int phase, TensorMap& env) +bool LanguageModel::Impl::has_embedding() const { - input_processor_->Run(BatchOp::kSetup, phase, env); - - auto& d = data_.at(phase); - auto& copy = *env.at("copy").data()[0]; - - Buffer_ rc = env.at("requests").buffer(); + return !weights_.decoder_only; +} - d.n_generating = 0; +bool LanguageModel::Impl::has_head() const +{ + return !weights_.decoder_only; +} - for (int i = 0; i < rc.size(); ++i) { - auto& c = *rc[i]; - d.autoregres[i] = c.autoregres; - d.generating[i] = c.generating; - d.n_generating += c.generating; - if (TM_UNLIKELY(!c.autoregres)) { - sequence_length_buf_[i] = c.history_len + c.inflight_input_len + c.input_len; - } - readonly_block_num_buf_[i] = c.readonly_block_num; // all rows, batch order +Tensor LanguageModel::Impl::Embed(const Buffer_& input_ids, Tensor out, const TensorMap& env) +{ + Buffer symm_buf; + if (comm_.d_comm) { + symm_buf = env.at("symm_buf").buffer(); } - - copy(sequence_length_buf_, rc.size(), d.sequence_length); - copy(readonly_block_num_buf_, rc.size(), d.readonly_block_num); - - unified_decoder_->Run(BatchOp::kSetup, phase, env); - generation_->Run(BatchOp::kSetup, phase, env); - output_processor_->Run(BatchOp::kSetup, phase, env); + return LookupEmbedding(input_ids, weights_.tok_embeddings, symm_buf, std::move(out)); } -void LanguageModel::Impl::Prepare(int phase, TensorMap& env) +LanguageModel::DecoderOutputs +LanguageModel::Impl::RunDecoder(int phase, const LanguageModel::DecoderInputs& in, TensorMap& env) { - env.emplace("autoreg_ids", autoreg_ids_); - - input_processor_->Run(BatchOp::kPrepare, phase, env); - - auto& d = data_.at(phase); - - auto& b = *env.at("batch").data()[0]; - auto& copy = *env.at("copy").data()[0]; + env.try_consume("hidden_states"); + env.try_consume("pre_final_residual"); + env.try_consume("full_hidden_states"); - // core::CopyT copy{}; + env.insert_or_assign("residual", in.residual); - if (auto group = copy.group()) { - for (int i = 0; i < b.bsz; ++i) { - if (const int j = b.perm[i]; j < b.bs0) { - copy(finished_.front().data() + j, 1, finished_.back().data() + i); - } - else { - copy(false_.data() + i, 1, finished_.back().data() + i); - } - } - finished_.Swap(); + if (in.attention_input) { + env.insert_or_assign("attention_input", in.attention_input); } - - if (auto group = copy.group()) { - // Non-autoregressive rows use the submitted prefix length: - // sequence_length = history_len + inflight_input_len + input_len. - // Existing autoregressive rows carry the previous sequence_length forward. - for (int i = 0; i < b.bsz; ++i) { - if (const int j = b.perm[i]; j < b.bs0 && d.autoregres[i]) { - copy(sequence_length_.front().data() + j, 1, sequence_length_.back().data() + i); - } - else { - copy(d.sequence_length.data() + i, 1, sequence_length_.back().data() + i); - } - } - sequence_length_.Swap(); + HiddenStateTap* tap = in.taps; + if (tap) { + env.insert_or_assign("hidden_state_tap", Buffer{&tap, 1, kCPU}); } - Buffer_ k_offsets{b.bsz + 1, kDEVICE}; - // PrefixSum(sequence_length_.front().data(), bsz, k_offsets.data(), core::Context::stream().handle()); - - // Buffer_ k_offsets_tmp{k_offsets.size(), kCPU}; - // Buffer_ sequence_length_tmp{sequence_length_.front().size(), kCPU}; + env.insert_or_assign("output_norm_weight", weights_.norm->weight); - // Copy(k_offsets, k_offsets_tmp); - // Copy(sequence_length_.front().buffer(), sequence_length_tmp); + env.insert_or_assign("selected_token_pos", in.selected_token_pos); - // core::Context::stream().Sync(); - - // dbg(core::to_vector(sequence_length_tmp.slice(0, bsz))); - // dbg(core::to_vector(k_offsets_tmp.slice(0, bsz + 1))); - - env.produce("finished", finished_.front()); - env.produce("sequence_length", sequence_length_.front()); - env.produce("readonly_block_num", d.readonly_block_num); - env.produce("k_offsets", k_offsets); - if (symm_buf_) { - env.produce("symm_buf", symm_buf_); + if (in.attention_metadata) { + unified_decoder_->SetAttentionForwardMetadata(phase, *in.attention_metadata); } - // Produced here so consumers may borrow the pointer at kPrepare; the content is - // only built at Forward time (`BuildTokenMask`). - env.produce("token_mask", token_mask_); + unified_decoder_->Forward(phase, env, weights_.layers_list(), in.selected_hidden_buffer); - unified_decoder_->Run(BatchOp::kPrepare, phase, env); - generation_->Run(BatchOp::kPrepare, phase, env); - output_processor_->Run(BatchOp::kPrepare, phase, env); + DecoderOutputs out; + out.selected_hidden = env.at("hidden_states"); + out.pre_final_residual = env.try_consume("pre_final_residual"); + return out; } -void LanguageModel::Impl::BuildTokenMask(const bool* finished, const int* q_offsets, const BatchData& b) +Tensor LanguageModel::Impl::Logits(const Tensor& hidden, Tensor out, const TensorMap& env) { - TM_FUNCTION_SCOPE(); - - if (b.global_token_num == 0) { - return; - } - - TM_CHECK_EQ((int)b.local_token_num.size(), attn_dp_size_); - TM_CHECK_LE(attn_dp_size_, kMaxAttnDPSize); - - const auto st = core::Context::stream().handle(); - - // Byte stride between per-rank metadata blocks (0 when attn_dp == 1). - size_t rank_stride = 0; - - if (attn_dp_size_ > 1) { - const int q_bytes = (max_batch_size_ + 1) * (int)sizeof(int); - const int meta_bytes = symm_token_meta_.shape(1); - - // Stage this rank's metadata into its row of the symmetric buffer; the finished - // tail is zeroed so padding slots never invalidate tokens. - TM_CHECK_LE(b.bsz, max_batch_size_); - char* slot = (char*)symm_token_meta_.data() + (ssize_t)attn_dp_rank_ * meta_bytes; - core::Copy(q_offsets, b.bsz + 1, (int*)slot); - core::Copy(finished, b.bsz, (bool*)(slot + q_bytes)); - TM_CUDA_CHECK(cudaMemsetAsync(slot + q_bytes + b.bsz, 0, max_batch_size_ - b.bsz, st)); - - // In-place all-gather: the peers read this rank's contribution from its own row. - comm_.d_comm->AllGather(slot, symm_token_meta_.data(), meta_bytes, kUint8, comm_.d_dp_group, st); - - q_offsets = (const int*)symm_token_meta_.data(); - finished = (const bool*)(symm_token_meta_.data() + q_bytes); - rank_stride = meta_bytes; + Buffer symm_buf; + if (comm_.d_comm) { + symm_buf = env.at("symm_buf").buffer(); } - - // Rank r's tokens occupy [token_base[r], token_base[r] + local_token_num[r]) of the mask. - int token_base[kMaxAttnDPSize]; - token_base[0] = 0; - std::partial_sum(b.local_token_num.begin(), b.local_token_num.end() - 1, token_base + 1); - - invokeBuildTokenMask(token_mask_.data(), - finished, - q_offsets, - rank_stride, - token_base, - attn_dp_size_, - // DP > 1 scans all gathered slots (the finished tail is zeroed); - // DP == 1 scans only the active batch — beyond it the local - // `finished`/`q_offsets` hold stale data from previous passes. - attn_dp_size_ > 1 ? max_batch_size_ : b.bsz, - b.global_token_num, - st); + return PostEmbedding(hidden, *weights_.output, symm_buf, std::move(out)); } -void LanguageModel::Impl::Forward(int phase, TensorMap& env) +const ModelWeight& LanguageModel::Impl::weights() const { - TM_FUNCTION_SCOPE(); - - auto& d = data_.at(phase); - auto& b = *env.at("batch").data()[0]; - - // Must run at Forward time: the `finished`/`q_offsets` H2D copies are only flushed - // after kPrepare returns. The mask is ready before the decoder (its consumers) runs. - BuildTokenMask( - (const bool*)env.at("finished").buffer().raw_data(), (const int*)env.at("q_offsets").buffer().raw_data(), b); - - { - Buffer_ k_offsets = env.at("k_offsets").buffer(); - PrefixSum(sequence_length_.front().data(), b.bsz, k_offsets.data(), core::Context::stream().handle()); - } - - { // compute input embeddings - auto input_ids = env.at("input_ids").buffer(); - - Tensor input_embeds = LookupEmbedding(input_ids, symm_buf_); - TM_DEBUG_TENSOR(input_embeds, "embeddings", 1); - - auto& copy = *env.at("copy").data()[0]; - input_processor_->PatchEmbedding(phase, input_embeds, copy, env); - copy.Run(); + return weights_; +} - env.produce("input_embeds", std::move(input_embeds)); - // dbg(env); +int LanguageModel::Impl::max_logits_len(const TensorMap& env) const +{ + if (has_head() && comm_.d_comm) { + return env.at("symm_buf").buffer().view(weights_.data_type).size() / weights_.vocab_size_padded; } + return max_logits_len_; +} - env.produce("output_norm_weight", weights_.norm->weight); - - unified_decoder_->Forward(phase, env, weights_.layers_list()); - - // env.at("batch").data()[0]->Notify(); +bool LanguageModel::Impl::logits_use_workspace() const +{ + return tp_size_ > 1 && use_ag2d_; +} - output_processor_->OutputHiddenStatesAndLogits(phase, env, 2); +void LanguageModel::Impl::CommitAcceptedState(int phase, const Buffer_& accept_len) +{ + unified_decoder_->CommitAcceptedState(phase, accept_len); +} - auto& hidden_states = env.at("hidden_states"); +size_t LanguageModel::Impl::SpeculativeStateJournalBytes(int request_count, int verification_positions) const +{ + return unified_decoder_->SpeculativeStateJournalBytes(request_count, verification_positions); +} - env.produce("logits", PostEmbedding(hidden_states, symm_buf_)); +LanguageModel::~LanguageModel() = default; - output_processor_->OutputHiddenStatesAndLogits(phase, env, 1); +LanguageModel::LanguageModel(LanguageModel&&) noexcept = default; - if (d.n_generating) { - generation_->Run(BatchOp::kForward, phase, env); - Copy(env.at("output_ids").buffer(), autoreg_ids_); - } +LanguageModel::LanguageModel(CacheRegistry& registry, + const EngineParam& engine, + const Context& ctx, + const ModelWeight& weights, + int phases) +{ + impl_ = std::make_unique(registry, engine, ctx, weights, phases); } -void LanguageModel::Impl::Unprep(int phase, TensorMap& env) +bool LanguageModel::has_embedding() const { - auto& d = data_.at(phase); - auto& copy = *env.at("copy").data()[0]; - - copy(sequence_length_.front().buffer(), d.sequence_length.size(), d.sequence_length); - - copy(finished_.front().buffer(), d.finished.size(), d.finished); + return TM_CHECK_NOTNULL(impl_)->has_embedding(); +} - unified_decoder_->Run(BatchOp::kUnprep, phase, env); - generation_->Run(BatchOp::kUnprep, phase, env); +bool LanguageModel::has_head() const +{ + return TM_CHECK_NOTNULL(impl_)->has_head(); } -void LanguageModel::Impl::Fetch(int phase, TensorMap& env) +Tensor LanguageModel::Embed(const Buffer_& input_ids, Tensor out, const TensorMap& env) { - auto& d = data_.at(phase); - auto& copy = *env.at("copy").data()[0]; + return TM_CHECK_NOTNULL(impl_)->Embed(input_ids, std::move(out), env); +} - copy(d.sequence_length, d.sequence_length.size(), sequence_length_buf_); - env.produce("sequence_length", sequence_length_buf_); +LanguageModel::DecoderOutputs +LanguageModel::RunDecoder(int phase, const DecoderInputs& in, TensorMap& env) +{ + return TM_CHECK_NOTNULL(impl_)->RunDecoder(phase, in, env); +} - copy(d.finished, d.finished.size(), finished_buf_); - env.produce("finished", finished_buf_); +Tensor LanguageModel::Logits(const Tensor& hidden, Tensor out, const TensorMap& env) +{ + return TM_CHECK_NOTNULL(impl_)->Logits(hidden, std::move(out), env); +} - env.produce("generating", d.generating); +const ModelWeight& LanguageModel::weights() const +{ + return TM_CHECK_NOTNULL(impl_)->weights(); +} - generation_->Run(BatchOp::kFetch, phase, env); +int LanguageModel::max_logits_len(const TensorMap& env) const +{ + return TM_CHECK_NOTNULL(impl_)->max_logits_len(env); } -LanguageModel::~LanguageModel() = default; +bool LanguageModel::logits_use_workspace() const +{ + return TM_CHECK_NOTNULL(impl_)->logits_use_workspace(); +} -LanguageModel::LanguageModel(LanguageModel&&) noexcept = default; +void LanguageModel::CommitAcceptedState(int phase, const Buffer_& accept_len) +{ + TM_CHECK_NOTNULL(impl_)->CommitAcceptedState(phase, accept_len); +} -LanguageModel::LanguageModel( - CacheRegistry& registry, const EngineParam& engine, const Context& ctx, const ModelWeight& weights, int phases) +size_t LanguageModel::SpeculativeStateJournalBytes(int request_count, int verification_positions) const { - impl_ = std::make_unique(registry, engine, ctx, weights, phases); + return TM_CHECK_NOTNULL(impl_)->SpeculativeStateJournalBytes(request_count, verification_positions); } void LanguageModel::Run(BatchOp op, int phase, TensorMap& env) diff --git a/src/turbomind/models/language_model.h b/src/turbomind/models/language_model.h index 0963b7755f..54cff048e0 100644 --- a/src/turbomind/models/language_model.h +++ b/src/turbomind/models/language_model.h @@ -10,6 +10,8 @@ namespace turbomind { class ModelWeight; +class HiddenStateTap; +struct AttentionForwardMetadata; struct Sequence; class CacheRegistry; @@ -34,6 +36,40 @@ class LanguageModel { void Run(BatchOp op, int phase, TensorMap& env); + bool has_embedding() const; + bool has_head() const; + + Tensor Embed(const Buffer_& input_ids, Tensor out, const TensorMap& env); + + struct DecoderInputs { + Tensor residual; + Tensor attention_input; + Buffer_ selected_token_pos; + Tensor selected_hidden_buffer; + + const AttentionForwardMetadata* attention_metadata{}; + HiddenStateTap* taps{}; + }; + + struct DecoderOutputs { + Tensor selected_hidden; + Tensor pre_final_residual; + }; + + DecoderOutputs RunDecoder(int phase, const DecoderInputs& in, TensorMap& env); + + Tensor Logits(const Tensor& hidden, Tensor out, const TensorMap& env); + + const ModelWeight& weights() const; + + int max_logits_len(const TensorMap& env) const; + + bool logits_use_workspace() const; + + void CommitAcceptedState(int phase, const Buffer_& accept_len); + + size_t SpeculativeStateJournalBytes(int request_count, int verification_positions) const; + private: struct Impl; std::unique_ptr impl_; diff --git a/src/turbomind/models/llama/GatedDeltaNetLayer.cc b/src/turbomind/models/llama/GatedDeltaNetLayer.cc index 5bb6ec36bf..d65339579f 100644 --- a/src/turbomind/models/llama/GatedDeltaNetLayer.cc +++ b/src/turbomind/models/llama/GatedDeltaNetLayer.cc @@ -9,10 +9,12 @@ #include "src/turbomind/core/allocator.h" #include "src/turbomind/core/check.h" +#include "src/turbomind/core/copy.h" #include "src/turbomind/core/data_type.h" #include "src/turbomind/core/logger.h" #include "src/turbomind/core/scope.h" #include "src/turbomind/engine/block.h" +#include "src/turbomind/kernels/copy/copy.h" #include "src/turbomind/models/llama/gated_delta_net_kernels.h" #include "src/turbomind/utils/cuda_utils.h" @@ -90,6 +92,8 @@ GatedDeltaNetLayer::GatedDeltaNetLayer(std::vector weights, num_v_heads_ = first.num_v_heads / tp_size_; head_dim_ = first.key_head_dim; gate_stride_ = num_v_heads_; + d_conv_ = first.d_conv; + conv_dim_ = 2 * num_k_heads_ * head_dim_ + num_v_heads_ * head_dim_; TM_CHECK_EQ(num_v_heads_ % num_k_heads_, 0); TM_CHECK(recurrent_state_dtype_ == kFloat32 || recurrent_state_dtype_ == input_dtype_) << "GDN recurrent state dtype must be float32 or match the input dtype, got state_dtype=" @@ -177,14 +181,26 @@ GatedDeltaNetLayer::GatedDeltaNetLayer(std::vector weights, layer_index_[weights[layer]] = layer; } - conv_state_ptrs_buf_ = {engine.max_batch_size, kCPUpinned}; - recurrent_state_ptrs_buf_ = {core::ssize_t(num_layer_groups_) * engine.max_batch_size * num_head_groups_, + conv_state_ptrs_buf_ = {engine.max_batch_size, kCPUpinned}; + recurrent_state_ptrs_buf_ = {core::ssize_t(num_layer_groups_) * engine.max_batch_size * num_head_groups_, kCPUpinned}; + speculative_request_indices_buf_ = {engine.max_batch_size, kCPUpinned}; + conv_state_offsets_buf_ = {layer_num_, kCPUpinned}; + for (int layer = 0; layer < layer_num_; ++layer) { + conv_state_offsets_buf_[layer] = static_cast(weights[layer]->conv_state_offset); + } for (int phase = 0; phase < phases; ++phase) { data_.emplace_back(); - data_.at(phase).conv_state_ptrs = empty_like(conv_state_ptrs_buf_, kDEVICE); - data_.at(phase).recurrent_state_ptrs = empty_like(recurrent_state_ptrs_buf_, kDEVICE); + data_.at(phase).conv_state_ptrs = empty_like(conv_state_ptrs_buf_, kDEVICE); + data_.at(phase).recurrent_state_ptrs = empty_like(recurrent_state_ptrs_buf_, kDEVICE); + data_.at(phase).state_store_suppressed = {engine.max_batch_size, kDEVICE}; + data_.at(phase).speculative_request_indices = {engine.max_batch_size, kDEVICE}; + data_.at(phase).conv_state_offsets = {layer_num_, kDEVICE}; + // Engine-thread kSetup has no pinned allocator. Keep each phase's + // staging storage alive until its executor-stream copy completes. + data_.at(phase).commit_state_ptrs_buf = { + core::ssize_t(layer_num_) * engine.max_batch_size * num_head_groups_, kCPUpinned}; } work_counter_ = {1, kDEVICE}; @@ -203,29 +219,87 @@ GatedDeltaNetLayer::~GatedDeltaNetLayer() void GatedDeltaNetLayer::Run(BatchOp op, int phase, TensorMap& env) { - if (op == BatchOp::kAdd) { - Buffer_ requests = env.at("requests").buffer(); - for (int i = 0; i < requests.size(); ++i) {} - } - else if (op == BatchOp::kSetup) { + if (op == BatchOp::kSetup) { Setup(phase, env); } else if (op == BatchOp::kPrepare) { - auto& data = data_.at(phase); - data.q_offsets = env.at("q_offsets").buffer().borrow(); - data.k_offsets = env.at("k_offsets").buffer().borrow(); - data.finished = env.at("finished").buffer().borrow(); + auto& data = data_.at(phase); + data.q_offsets = env.at("q_offsets").buffer().borrow(); + data.k_offsets = env.at("k_offsets").buffer().borrow(); + data.entry_sequence_length = env.at("sequence_length").buffer().borrow(); + // Verification rows in the batch — counted at kSetup from the + // submitted rows — select the store-suppressed path. The + // speculative-row flag itself flows to the mask kernel as data. + if (data.speculative_request_count != 0) { + data.finished_on_entry = env.at("finished").buffer().borrow(); + data.speculative_row = env.at("speculative_row").buffer().borrow(); + data.finished = data.state_store_suppressed.slice(0, data.batch_size); + data.build_state_store_mask = true; + } + else { + data.finished = env.at("finished").buffer().borrow(); + data.finished_on_entry = {}; + data.speculative_row = {}; + data.build_state_store_mask = false; + } + + if (data.speculative_request_count != 0) { + const int L = layer_num_; + const int Sspec = data.speculative_request_count; + const int K = data.verify_positions; + data.journal.raw_conv = {{L, Sspec, K, conv_dim_}, input_dtype_, kDEVICE}; + data.journal.key = {{L, Sspec, K, num_k_heads_, 128}, input_dtype_, kDEVICE}; + data.journal.value = {{L, Sspec, K, num_v_heads_, 128}, input_dtype_, kDEVICE}; + data.journal.log_decay = {{L, Sspec, K, num_v_heads_}, kFloat32, kDEVICE}; + data.journal.beta = {{L, Sspec, K, num_v_heads_}, kFloat32, kDEVICE}; + } for (const auto& [ptr, bytes] : data.reset_ptrs) { Clear(Buffer_{ptr, static_cast(bytes), kDEVICE}); } data.reset_ptrs.clear(); - if (data.recurrent_plan) { + if (data.commit_plan) { + const core::ssize_t count = core::ssize_t(layer_num_) * data.speculative_request_count; + auto host_ptrs = data.commit_state_ptrs_buf.slice(0, count * num_head_groups_); + data.commit_state_ptrs = Tensor{empty_like(host_ptrs, kDEVICE), + core::Layout{{count, num_head_groups_}}}; + data.commit_state_tma_descs = {{count, num_head_groups_, 128}, kUint8, kDEVICE}; + data.commit_lengths = {{count}, kInt32, kDEVICE}; + Copy(host_ptrs, data.commit_state_ptrs.buffer()); + + // Each flattened request points at one layer slice of its cache part. + auto state_ptrs = data.commit_state_ptrs.view({1, count, num_head_groups_}); + auto state_descs = data.commit_state_tma_descs.view({1, count, num_head_groups_, 128}); + delta_rule_.PrepareState(state_ptrs, + state_descs, + 1, + 1, + *data.commit_plan, + core::Context::stream().handle()); + } + + if (data.verify_plan) { core::Tensor state_ptrs{data.recurrent_state_ptrs, + core::Layout{{num_layer_groups_, data.verify_count, num_head_groups_}, + {data.batch_size * num_head_groups_, num_head_groups_, 1}}, + core::Tensor::PreserveBufferCapacity{}}; + core::Tensor state_descs{data.verify_state_tma_descs, + core::Layout{{num_layer_groups_, data.verify_count, num_head_groups_, 128}}}; + delta_rule_.PrepareState(state_ptrs, + state_descs, + num_layer_groups_, + layers_per_block_, + *data.verify_plan, + core::Context::stream().handle()); + } + if (data.recurrent_plan) { + const core::ssize_t base = core::ssize_t(data.verify_count) * num_head_groups_; + const auto tail = data.recurrent_state_ptrs.slice(base, data.recurrent_state_ptrs.size() - base); + core::Tensor state_ptrs{tail, core::Layout{{num_layer_groups_, data.decode_count, num_head_groups_}, {data.batch_size * num_head_groups_, num_head_groups_, 1}}, core::Tensor::PreserveBufferCapacity{}}; - core::Tensor state_descs; + core::Tensor state_descs; if (data.recurrent_state_tma_descs) { state_descs = core::Tensor{data.recurrent_state_tma_descs, core::Layout{{num_layer_groups_, data.decode_count, num_head_groups_, 128}}}; @@ -248,23 +322,30 @@ void GatedDeltaNetLayer::Setup(int phase, TensorMap& env) data.batch_size = requests.size(); data.input_lens.resize(data.batch_size); data.reset_ptrs.clear(); + data.speculative_request_count = 0; + data.verify_positions = env.at("verification_positions").data()[0]; std::vector host_offsets(data.batch_size + 1, 0); for (int sequence = 0; sequence < data.batch_size; ++sequence) { - data.input_lens[sequence] = requests[sequence]->input_len; + const Sequence& request = *requests[sequence]; + const SubmittedRow& row = *request.submitted; + data.input_lens[sequence] = row.input_len; + if (row.is_verification_row()) { + speculative_request_indices_buf_[data.speculative_request_count++] = sequence; + } host_offsets[sequence + 1] = host_offsets[sequence] + data.input_lens[sequence]; } - const int token_slots = *env.at("token_num").data(); - data.decode_count = 0; - while (data.decode_count < data.batch_size && data.input_lens[data.decode_count] == 1) { - ++data.decode_count; - } - data.prefill_count = data.batch_size - data.decode_count; - + const int token_slots = *env.at("token_num").data(); + data.verify_count = 0; + data.decode_count = 0; + data.prefill_count = 0; + data.verify_plan.reset(); + data.commit_plan.reset(); data.recurrent_plan.reset(); data.chunked_plan.reset(); - data.chunked_workspace = {}; - data.recurrent_state_tma_descs = {}; + data.chunked_workspace = {}; + data.verify_state_tma_descs = {}; + data.recurrent_state_tma_descs = {}; auto make_context = [&] { linear_attn::delta_rule::PlanningContext planning{}; @@ -282,6 +363,52 @@ void GatedDeltaNetLayer::Setup(int phase, TensorMap& env) return planning; }; + const bool verify_candidate = arch_ == 900 && input_dtype_ == kBfloat16 && head_dim_ == 128 + && data.speculative_request_count != 0 && data.verify_positions <= 16; + if (verify_candidate) { + auto planning = make_context(); + planning.physical_batch = data.speculative_request_count; + planning.token_slots = data.verify_positions; + planning.gate_batch_stride = int64_t(data.verify_positions) * gate_stride_; + planning.beta_batch_stride = planning.gate_batch_stride; + linear_attn::delta_rule::Operation operation{}; + operation.mode = linear_attn::delta_rule::GdrMode::kVerify; + operation.chunk_size = data.verify_positions <= 8 ? 8 : 16; + operation.cp_level = ContextParallelLevel::kOff; + linear_attn::delta_rule::Plan plan; + if (delta_rule_.Plan(operation, planning, &plan)) { + data.verify_plan.emplace(std::move(plan)); + data.verify_count = data.speculative_request_count; + } + } + + const bool commit_candidate = arch_ == 900 && input_dtype_ == kBfloat16 && head_dim_ == 128 + && data.speculative_request_count != 0 + && data.verify_positions >= 1 && data.verify_positions <= 16; + if (commit_candidate) { + auto planning = make_context(); + planning.physical_batch = layer_num_ * data.speculative_request_count; + planning.token_slots = data.verify_positions; + planning.gate_stride = num_v_heads_; + planning.beta_stride = num_v_heads_; + planning.gate_batch_stride = int64_t(data.verify_positions) * num_v_heads_; + planning.beta_batch_stride = planning.gate_batch_stride; + linear_attn::delta_rule::Operation operation{}; + operation.mode = linear_attn::delta_rule::GdrMode::kCommit; + operation.chunk_size = data.verify_positions <= 8 ? 8 : 16; + operation.cp_level = ContextParallelLevel::kOff; + linear_attn::delta_rule::Plan plan; + if (delta_rule_.Plan(operation, planning, &plan)) { + data.commit_plan.emplace(std::move(plan)); + } + } + + while (data.verify_count + data.decode_count < data.batch_size + && data.input_lens[data.verify_count + data.decode_count] == 1) { + ++data.decode_count; + } + data.prefill_count = data.batch_size - data.verify_count - data.decode_count; + if (data.decode_count != 0) { auto planning = make_context(); planning.physical_batch = data.decode_count; @@ -301,7 +428,8 @@ void GatedDeltaNetLayer::Setup(int phase, TensorMap& env) planning.token_slots = token_slots; planning.gate_batch_stride = int64_t(token_slots) * gate_stride_; planning.beta_batch_stride = planning.gate_batch_stride; - planning.q_offsets.assign(host_offsets.begin() + data.decode_count, host_offsets.end()); + const int first_prefill = data.verify_count + data.decode_count; + planning.q_offsets.assign(host_offsets.begin() + first_prefill, host_offsets.end()); linear_attn::delta_rule::Operation operation{}; operation.mode = linear_attn::delta_rule::GdrMode::kChunked; operation.cp_level = gdr_cp_level_; @@ -314,6 +442,11 @@ void GatedDeltaNetLayer::Setup(int phase, TensorMap& env) data.chunked_workspace = core::Tensor{ core::Layout{{static_cast(data.chunked_plan->workspace_bytes)}}, kUint8, kDEVICE}; } + if (data.verify_plan && data.verify_plan->state_tma_desc_bytes_per_layer_group != 0) { + const core::ssize_t descriptor_bytes = + core::ssize_t(num_layer_groups_) * data.verify_plan->state_tma_desc_bytes_per_layer_group; + data.verify_state_tma_descs = {descriptor_bytes, kDEVICE}; + } if (data.recurrent_plan && data.recurrent_plan->state_tma_desc_bytes_per_layer_group != 0) { const core::ssize_t descriptor_bytes = core::ssize_t(num_layer_groups_) * data.recurrent_plan->state_tma_desc_bytes_per_layer_group; @@ -321,7 +454,8 @@ void GatedDeltaNetLayer::Setup(int phase, TensorMap& env) } for (int sequence = 0; sequence < data.batch_size; ++sequence) { - auto& request = *requests[sequence]; + auto& request = *requests[sequence]; + const SubmittedRow& row = *request.submitted; const CacheBlock& block = *TM_CHECK_NOTNULL(request.frontier.get()); TM_CHECK_NOTNULL(block.allocation.a); @@ -335,7 +469,7 @@ void GatedDeltaNetLayer::Setup(int phase, TensorMap& env) } } - if (request.history_len + request.inflight_input_len == 0) { + if (row.history_len + request.inflight_input_len == 0) { data.reset_ptrs.push_back({reinterpret_cast(block.base(0)), conv_total_bytes_}); for (int recurrent_block = 0; recurrent_block < num_blocks_; ++recurrent_block) { data.reset_ptrs.push_back( @@ -344,10 +478,37 @@ void GatedDeltaNetLayer::Setup(int phase, TensorMap& env) } } + if (data.commit_plan) { + const int requests = data.speculative_request_count; + // contracts.scheduler-output puts speculative rows first in executor order. + for (int request = 0; request < requests; ++request) { + TM_CHECK_EQ(speculative_request_indices_buf_[request], request); + } + for (int layer = 0; layer < layer_num_; ++layer) { + const int layer_group = layer / layers_per_block_; + const size_t layer_bytes = byte_size(recurrent_state_dtype_, + size_t(layer % layers_per_block_) * heads_per_block_ * head_dim_ * head_dim_); + for (int request = 0; request < requests; ++request) { + for (int head_group = 0; head_group < num_head_groups_; ++head_group) { + const core::ssize_t source = + (core::ssize_t(layer_group) * data.batch_size + request) * num_head_groups_ + head_group; + const core::ssize_t destination = + (core::ssize_t(layer) * requests + request) * num_head_groups_ + head_group; + data.commit_state_ptrs_buf[destination] = + static_cast(recurrent_state_ptrs_buf_[source]) + layer_bytes; + } + } + } + } + Copy(conv_state_ptrs_buf_, data.batch_size, data.conv_state_ptrs); Copy(recurrent_state_ptrs_buf_, core::ssize_t(num_layer_groups_) * data.batch_size * num_head_groups_, data.recurrent_state_ptrs); + + auto& copy = *env.at("copy").data()[0]; + copy(speculative_request_indices_buf_, data.speculative_request_count, data.speculative_request_indices); + copy(conv_state_offsets_buf_, layer_num_, data.conv_state_offsets); } void GatedDeltaNetLayer::Forward(ForwardParam param) @@ -364,6 +525,15 @@ void GatedDeltaNetLayer::Forward(ForwardParam param) const auto stream = core::Context::stream().handle(); const auto& weights = *param.weights; auto& phase_data = data_.at(param.phase); + const int layer = layer_index_.at(param.weights); + + if (phase_data.build_state_store_mask && layer == 0) { + linear_attn::delta_rule::invokeBuildGdnStateStoreMask(phase_data.state_store_suppressed.data(), + phase_data.finished_on_entry.data(), + phase_data.speculative_row.data(), + phase_data.batch_size, + stream); + } TM_CHECK(dtype == kHalf || dtype == kBfloat16); @@ -420,7 +590,21 @@ void GatedDeltaNetLayer::Forward(ForwardParam param) Tensor out{attn_out.buffer(), out_layout}; invokeL2NormalizeQK(q, k, 1e-6f, stream); - const int layer = layer_index_.at(param.weights); + if (phase_data.speculative_request_count != 0) { + linear_attn::delta_rule::invokeCaptureGdnTransitions( + all_proj.slice({0, 0}, {-1, conv_dim}), + k, + v, + g, + beta, + phase_data.q_offsets, + phase_data.speculative_request_indices.slice(0, phase_data.speculative_request_count), + layer, + phase_data.verify_positions, + phase_data.journal, + stream); + } + const int layer_group = layer / layers_per_block_; const int64_t state_layer_offset = weights.linear_state_offset; @@ -432,59 +616,88 @@ void GatedDeltaNetLayer::Forward(ForwardParam param) core::Layout{{sequence_count, num_head_groups_}}}; }; - const bool mixed = phase_data.recurrent_plan.has_value() && phase_data.chunked_plan.has_value(); - if (mixed) { - TM_CUDA_CHECK(cudaEventRecord(ev_before_, stream)); - TM_CUDA_CHECK(cudaStreamWaitEvent(aux_stream_, ev_before_)); + const int S = phase_data.verify_count; + const int K = phase_data.verify_positions; + + linear_attn::delta_rule::Arguments verify_args{}; + Tensor verify_out; + if (phase_data.verify_plan) { + const core::Layout verify_qk_layout{{S, K, num_k_heads_, 128}, {int64_t(K) * conv_dim, conv_dim, 128, 1}}; + const core::Layout verify_v_layout{{S, K, num_v_heads_, 128}, {int64_t(K) * conv_dim, conv_dim, 128, 1}}; + const core::Layout verify_gate_layout{{S, K, num_v_heads_}, {int64_t(K) * gate_stride_, gate_stride_, 1}}; + const core::Layout verify_out_layout{{S, K, num_v_heads_, 128}, {int64_t(K) * value_dim, value_dim, 128, 1}}; + verify_out = Tensor{out.buffer(), verify_out_layout, Tensor::PreserveBufferCapacity{}}; + verify_args.q = Tensor{q.buffer(), verify_qk_layout, Tensor::PreserveBufferCapacity{}}; + verify_args.k = Tensor{k.buffer(), verify_qk_layout, Tensor::PreserveBufferCapacity{}}; + verify_args.v = Tensor{v.buffer(), verify_v_layout, Tensor::PreserveBufferCapacity{}}; + verify_args.g = Tensor{g.buffer(), verify_gate_layout, Tensor::PreserveBufferCapacity{}}; + verify_args.beta = Tensor{beta.buffer(), verify_gate_layout, Tensor::PreserveBufferCapacity{}}; + verify_args.state_ptrs = pointer_view(0, S); + const core::ssize_t descriptor_count = core::ssize_t(S) * num_head_groups_ * 128; + const core::ssize_t descriptor_offset = core::ssize_t(layer_group) * descriptor_count; + verify_args.state_tma_descs = + Tensor{phase_data.verify_state_tma_descs.slice(descriptor_offset, descriptor_count), + core::Layout{{S, num_head_groups_, 128}}}; + verify_args.out = &verify_out; + verify_args.state_layer_offset = state_layer_offset; } - const cudaStream_t chunk_stream = mixed ? aux_stream_ : stream; + linear_attn::delta_rule::Arguments recurrent_args{}; + Tensor recurrent_out; if (phase_data.recurrent_plan) { + const int ordinary_token_begin = S * K; + auto token_tail = [ordinary_token_begin](const Tensor& tensor) { + const core::ssize_t first = core::ssize_t(ordinary_token_begin) * tensor.stride(1); + return tensor.buffer().slice(first, tensor.buffer().size() - first); + }; const core::Layout recurrent_qk_layout{{phase_data.decode_count, 1, num_k_heads_, 128}, {conv_dim, conv_dim, 128, 1}}; const core::Layout recurrent_v_layout{{phase_data.decode_count, 1, num_v_heads_, 128}, {conv_dim, conv_dim, 128, 1}}; - const core::Layout recurrent_out_layout{{phase_data.decode_count, 1, num_v_heads_, 128}, - {value_dim, value_dim, 128, 1}}; const core::Layout recurrent_gate_layout{{phase_data.decode_count, 1, num_v_heads_}, {gate_stride_, gate_stride_, 1}}; - Tensor recurrent_q{q.buffer(), recurrent_qk_layout, Tensor::PreserveBufferCapacity{}}; - Tensor recurrent_k{k.buffer(), recurrent_qk_layout, Tensor::PreserveBufferCapacity{}}; - Tensor recurrent_v{v.buffer(), recurrent_v_layout, Tensor::PreserveBufferCapacity{}}; - Tensor recurrent_out{out.buffer(), recurrent_out_layout, Tensor::PreserveBufferCapacity{}}; - Tensor recurrent_g{g.buffer(), recurrent_gate_layout, Tensor::PreserveBufferCapacity{}}; - Tensor recurrent_beta{beta.buffer(), recurrent_gate_layout, Tensor::PreserveBufferCapacity{}}; - Tensor recurrent_state_ptrs = pointer_view(0, phase_data.decode_count); - Tensor recurrent_finished{phase_data.finished.slice(0, phase_data.decode_count), - core::Layout{{phase_data.decode_count}}}; - Tensor recurrent_state_descs; + const core::Layout recurrent_out_layout{{phase_data.decode_count, 1, num_v_heads_, 128}, + {value_dim, value_dim, 128, 1}}; + recurrent_out = Tensor{token_tail(out), recurrent_out_layout, Tensor::PreserveBufferCapacity{}}; + recurrent_args.q = Tensor{token_tail(q), recurrent_qk_layout, Tensor::PreserveBufferCapacity{}}; + recurrent_args.k = Tensor{token_tail(k), recurrent_qk_layout, Tensor::PreserveBufferCapacity{}}; + recurrent_args.v = Tensor{token_tail(v), recurrent_v_layout, Tensor::PreserveBufferCapacity{}}; + recurrent_args.g = Tensor{token_tail(g), recurrent_gate_layout, Tensor::PreserveBufferCapacity{}}; + recurrent_args.beta = Tensor{token_tail(beta), recurrent_gate_layout, Tensor::PreserveBufferCapacity{}}; + recurrent_args.state_ptrs = pointer_view(S, phase_data.decode_count); if (phase_data.recurrent_state_tma_descs) { const core::ssize_t descriptor_count = core::ssize_t(phase_data.decode_count) * num_head_groups_ * 128; const core::ssize_t descriptor_offset = core::ssize_t(layer_group) * descriptor_count; - recurrent_state_descs = + recurrent_args.state_tma_descs = Tensor{phase_data.recurrent_state_tma_descs.slice(descriptor_offset, descriptor_count), core::Layout{{phase_data.decode_count, num_head_groups_, 128}}}; } + recurrent_args.finished = + Tensor{phase_data.finished.slice(S, phase_data.decode_count), core::Layout{{phase_data.decode_count}}}; + recurrent_args.out = &recurrent_out; + recurrent_args.state_layer_offset = state_layer_offset; + } - linear_attn::delta_rule::Arguments arguments{}; - arguments.q = recurrent_q; - arguments.k = recurrent_k; - arguments.v = recurrent_v; - arguments.g = recurrent_g; - arguments.beta = recurrent_beta; - arguments.state_ptrs = recurrent_state_ptrs; - arguments.state_tma_descs = recurrent_state_descs; - arguments.finished = recurrent_finished; - arguments.out = &recurrent_out; - arguments.state_layer_offset = state_layer_offset; - delta_rule_.Run(arguments, *phase_data.recurrent_plan, stream); + const bool has_main = phase_data.verify_plan.has_value() || phase_data.recurrent_plan.has_value(); + const bool fork_chunked = has_main && phase_data.chunked_plan.has_value(); + if (fork_chunked) { + TM_CUDA_CHECK(cudaEventRecord(ev_before_, stream)); + TM_CUDA_CHECK(cudaStreamWaitEvent(aux_stream_, ev_before_)); + } + + if (phase_data.verify_plan) { + delta_rule_.Run(verify_args, *phase_data.verify_plan, stream); + } + if (phase_data.recurrent_plan) { + delta_rule_.Run(recurrent_args, *phase_data.recurrent_plan, stream); } if (phase_data.chunked_plan) { - Tensor chunk_state_ptrs = pointer_view(phase_data.decode_count, phase_data.prefill_count); - Tensor chunk_finished{phase_data.finished.slice(phase_data.decode_count, phase_data.prefill_count), + const int first_prefill = S + phase_data.decode_count; + Tensor chunk_state_ptrs = pointer_view(first_prefill, phase_data.prefill_count); + Tensor chunk_finished{phase_data.finished.slice(first_prefill, phase_data.prefill_count), core::Layout{{phase_data.prefill_count}}}; - Tensor chunk_q_offsets{phase_data.q_offsets.slice(phase_data.decode_count, phase_data.prefill_count + 1), + Tensor chunk_q_offsets{phase_data.q_offsets.slice(first_prefill, phase_data.prefill_count + 1), core::Layout{{phase_data.prefill_count + 1}}}; linear_attn::delta_rule::Arguments arguments{}; @@ -499,10 +712,10 @@ void GatedDeltaNetLayer::Forward(ForwardParam param) arguments.out = &out; arguments.workspace = phase_data.chunked_workspace ? &phase_data.chunked_workspace : nullptr; arguments.state_layer_offset = state_layer_offset; - delta_rule_.Run(arguments, *phase_data.chunked_plan, chunk_stream); + delta_rule_.Run(arguments, *phase_data.chunked_plan, fork_chunked ? aux_stream_ : stream); } - if (mixed) { + if (fork_chunked) { TM_CUDA_CHECK(cudaEventRecord(ev_after_, aux_stream_)); TM_CUDA_CHECK(cudaStreamWaitEvent(stream, ev_after_)); } @@ -514,4 +727,96 @@ void GatedDeltaNetLayer::Forward(ForwardParam param) TM_SCOPE_CALL(linear_.Forward(attn_out, *weights.out_proj, param.output)); } +void GatedDeltaNetLayer::CommitAcceptedState(int phase, const Buffer_& accept_len) +{ + Data& data = data_.at(phase); + if (data.speculative_request_count == 0) { + return; + } + + const cudaStream_t stream = core::Context::stream().handle(); + const auto request_indices = data.speculative_request_indices.slice(0, data.speculative_request_count); + + linear_attn::delta_rule::invokeCommitAcceptedConvState(data.journal.raw_conv, + data.conv_state_ptrs, + request_indices, + data.entry_sequence_length, + accept_len, + data.conv_state_offsets, + conv_dim_, + d_conv_, + stream); + + if (data.commit_plan) { + const int requests = data.speculative_request_count; + const int positions = data.verify_positions; + const core::ssize_t count = core::ssize_t(layer_num_) * requests; + Tensor accepted_lengths{accept_len.slice(0, requests), + core::Layout{{layer_num_, requests}, {0, 1}}}; + auto commit_lengths = data.commit_lengths.view({layer_num_, requests}); + core::GenericCopy(accepted_lengths, commit_lengths, stream); + + linear_attn::delta_rule::Arguments args{}; + args.k = data.journal.key.view({count, positions, num_k_heads_, 128}); + args.v = data.journal.value.view({count, positions, num_v_heads_, 128}); + args.g = data.journal.log_decay.view({count, positions, num_v_heads_}); + args.beta = data.journal.beta.view({count, positions, num_v_heads_}); + args.state_ptrs = data.commit_state_ptrs; + args.state_tma_descs = data.commit_state_tma_descs; + args.commit_lengths = data.commit_lengths; + // Final accept_len is already terminal-clamped. Newly finished rows + // still commit their accepted prefix; no forward suppression mask applies. + // Adjusted pointer bases make every logical request a single-layer state. + args.state_layer_offset = 0; + delta_rule_.Run(args, *data.commit_plan, stream); + } + else { + Tensor recurrent_state_ptrs{data.recurrent_state_ptrs, + core::Layout{{num_layer_groups_, data.batch_size, num_head_groups_}, + {data.batch_size * num_head_groups_, num_head_groups_, 1}}, + Tensor::PreserveBufferCapacity{}}; + + linear_attn::delta_rule::AcceptedPrefixArguments args{}; + args.key = data.journal.key; + args.value = data.journal.value; + args.log_decay = data.journal.log_decay; + args.beta = data.journal.beta; + args.recurrent_state_ptrs = recurrent_state_ptrs; + args.request_indices = Tensor{request_indices, core::Layout{{data.speculative_request_count}}}; + args.accept_len = Tensor{accept_len, core::Layout{{data.batch_size}}}; + args.layer_count = layer_num_; + args.speculative_count = data.speculative_request_count; + args.position_count = data.verify_positions; + args.hq = num_k_heads_; + args.hv = num_v_heads_; + args.num_head_groups = num_head_groups_; + args.layers_per_block = layers_per_block_; + args.heads_per_block = heads_per_block_; + args.sm_count = sm_count_; + delta_rule_.CommitAccepted(args, recurrent_state_dtype_, stream); + } + + // These buffers were allocated during kPrepare on the executor stream. + // Release them with the journal after their last queued consumers. + data.journal = {}; + data.commit_state_ptrs = {}; + data.commit_state_tma_descs = {}; + data.commit_lengths = {}; +} + +size_t GatedDeltaNetLayer::SpeculativeStateJournalBytes(int request_count, int verify_positions) const +{ + const size_t rows = size_t(layer_num_) * request_count * verify_positions; + const size_t input_elements = size_t(conv_dim_) + size_t(num_k_heads_ + num_v_heads_) * 128; + const size_t gate_elements = size_t(2) * num_v_heads_; + size_t bytes = rows * (byte_size(input_dtype_, input_elements) + gate_elements * sizeof(float)); + if (arch_ == 900 && input_dtype_ == kBfloat16 && head_dim_ == 128 + && verify_positions >= 1 && verify_positions <= 16) { + // Commit metadata has the same executor-stream lifetime as the journal. + const size_t commit_rows = size_t(layer_num_) * request_count; + bytes += commit_rows * (num_head_groups_ * (sizeof(void*) + 128) + sizeof(int)); + } + return bytes; +} + } // namespace turbomind diff --git a/src/turbomind/models/llama/GatedDeltaNetLayer.h b/src/turbomind/models/llama/GatedDeltaNetLayer.h index dc0464f4df..32fd7ecc0d 100644 --- a/src/turbomind/models/llama/GatedDeltaNetLayer.h +++ b/src/turbomind/models/llama/GatedDeltaNetLayer.h @@ -38,6 +38,10 @@ class GatedDeltaNetLayer { void Forward(ForwardParam p); + void CommitAcceptedState(int phase, const Buffer_& accept_len); + + size_t SpeculativeStateJournalBytes(int request_count, int verify_positions) const; + private: void Setup(int phase, TensorMap& env); @@ -57,13 +61,31 @@ class GatedDeltaNetLayer { Buffer_ q_offsets; Buffer_ k_offsets; Buffer_ finished; + Buffer_ finished_on_entry; + Buffer_ speculative_row; + Buffer_ state_store_suppressed; + bool build_state_store_mask{}; + Buffer_ speculative_request_indices; + Buffer_ conv_state_offsets; + int speculative_request_count{}; + int verify_positions{}; + Buffer_ entry_sequence_length; + linear_attn::delta_rule::TransitionJournal journal; Buffer_ conv_state_ptrs; Buffer_ recurrent_state_ptrs; + int verify_count{}; int decode_count{}; int prefill_count{}; + std::optional verify_plan; + std::optional commit_plan; std::optional recurrent_plan; std::optional chunked_plan; + Buffer_ commit_state_ptrs_buf; + core::Tensor commit_state_ptrs; + core::Tensor commit_state_tma_descs; + core::Tensor commit_lengths; core::Tensor chunked_workspace; + Buffer_ verify_state_tma_descs; Buffer_ recurrent_state_tma_descs; }; std::vector data_; @@ -83,6 +105,8 @@ class GatedDeltaNetLayer { // staging buffers Buffer_ conv_state_ptrs_buf_; Buffer_ recurrent_state_ptrs_buf_; + Buffer_ speculative_request_indices_buf_; + Buffer_ conv_state_offsets_buf_; DataType input_dtype_{kNull}; int arch_{}; @@ -90,6 +114,8 @@ class GatedDeltaNetLayer { int num_v_heads_{}; int head_dim_{}; int gate_stride_{}; + int d_conv_{}; + int conv_dim_{}; linear_attn::delta_rule::GatedDeltaRule delta_rule_; int sm_count_{}; diff --git a/src/turbomind/models/llama/LlamaFfnLayer.cc b/src/turbomind/models/llama/LlamaFfnLayer.cc index 406eb546b4..f0ed051bf1 100644 --- a/src/turbomind/models/llama/LlamaFfnLayer.cc +++ b/src/turbomind/models/llama/LlamaFfnLayer.cc @@ -22,6 +22,7 @@ #include "src/turbomind/kernels/activation.h" #include "src/turbomind/models/llama/llama_utils.h" #include "src/turbomind/utils/anomaly_handler.h" +#include "src/turbomind/utils/nvtx_utils.h" namespace turbomind { diff --git a/src/turbomind/models/llama/context_token_resource.h b/src/turbomind/models/llama/context_token_resource.h index 292eb8d64c..21ca6ae33c 100644 --- a/src/turbomind/models/llama/context_token_resource.h +++ b/src/turbomind/models/llama/context_token_resource.h @@ -10,22 +10,18 @@ class ContextTokenResource final: public Resource { public: explicit ContextTokenResource(int max_context_tokens) noexcept: max_context_tokens_{max_context_tokens} {} - int Test(const Sequence& s) const noexcept override + int Test(const Sequence& s, const SubmittedRow& row) const noexcept override { - const int input_len = InputLen(s, s.resume_len); - if (input_len <= 0) { + const int q = row.query_count; + if (q <= 0) { return 0; } - if (TempLen(s, input_len) > max_context_tokens_) { - return 0; - } - return input_len; + return Charge(s, row) <= max_context_tokens_ ? q : 0; } - void Commit(const Sequence& s) noexcept override + void Commit(const Sequence& s, const SubmittedRow& row) noexcept override { - const int input_len = InputLen(s, s.history_len); - max_context_tokens_ -= TempLen(s, input_len); + max_context_tokens_ -= Charge(s, row); } int remaining_tokens() const noexcept @@ -39,14 +35,15 @@ class ContextTokenResource final: public Resource { return s.seq_len + s.inflight_new_tokens; } - static int InputLen(const Sequence& s, int history_len) noexcept + static int Charge(const Sequence& s, const SubmittedRow& row) noexcept { - return ContextLen(s) - s.inflight_input_len - history_len; - } + if (row.is_verification_row()) { + return row.key_capacity_end; + } - static int TempLen(const Sequence& s, int input_len) noexcept - { - return (input_len > 1 || !s.is_active) ? ContextLen(s) : 0; + const int context_len = ContextLen(s); + const int remaining_input = context_len - s.inflight_input_len - row.history_len; + return (remaining_input > 1 || !s.is_active) ? context_len : 0; } int max_context_tokens_{}; diff --git a/src/turbomind/models/llama/llama_kernels.cu b/src/turbomind/models/llama/llama_kernels.cu index 7009818bc9..44f47335c6 100644 --- a/src/turbomind/models/llama/llama_kernels.cu +++ b/src/turbomind/models/llama/llama_kernels.cu @@ -469,6 +469,10 @@ __global__ void CollectHiddenStates_Kernel(const T* src, const int* idxs, T* dst void CollectHiddenStates(const Tensor& src, const Buffer_& idxs, Ref dst, cudaStream_t st) { + if (idxs.size() == 0) { + return; + } + const auto stride = byte_size(src.dtype(), src.stride(0)); auto invoke = [&](auto t) { @@ -493,9 +497,8 @@ void CollectHiddenStates(const Tensor& src, const Buffer_& idxs, Ref @@ -553,23 +556,32 @@ void BatchPrefixSum(const int** srcs, const int* ns, int** dsts, int count, cuda TM_CUDA_CHECK(cudaGetLastError()); } -__global__ void AppendTokenIdsKernel(int** token_ids_ptrs, const int* output_ids, const int* positions, int batch_size) +__global__ void AppendOneTokenAndAdvanceSequenceKernel(int* const* token_ids_ptrs, + const int* selected_tokens, + int* sequence_length, + int batch_size) { - int i = threadIdx.x + blockIdx.x * blockDim.x; - if (i < batch_size) { - int* token_ids = token_ids_ptrs[i]; - int pos = positions[i]; - token_ids[pos] = output_ids[i]; + const int b = blockIdx.x * blockDim.x + threadIdx.x; + + if (b < batch_size) { + const int position = sequence_length[b]; + token_ids_ptrs[b][position] = selected_tokens[b]; + sequence_length[b] = position + 1; } } -void AppendTokenIds( - int** token_ids_ptrs, const int* output_ids, const int* positions, int batch_size, cudaStream_t stream) +void invokeAppendOneTokenAndAdvanceSequence( + int* const* token_ids_ptrs, const int* selected_tokens, int* sequence_length, int batch_size, cudaStream_t stream) { - constexpr int block = 128; - const int grid = cdiv(batch_size, block); - AppendTokenIdsKernel<<>>(token_ids_ptrs, output_ids, positions, batch_size); - TM_CUDA_CHECK(cudaGetLastError()); + if (batch_size == 0) { + return; + } + + constexpr int block_size = 128; + const int grid_size = cdiv(batch_size, block_size); + + AppendOneTokenAndAdvanceSequenceKernel<<>>( + token_ids_ptrs, selected_tokens, sequence_length, batch_size); } template diff --git a/src/turbomind/models/llama/llama_kernels.h b/src/turbomind/models/llama/llama_kernels.h index 9dadeb9d89..783c34b09f 100644 --- a/src/turbomind/models/llama/llama_kernels.h +++ b/src/turbomind/models/llama/llama_kernels.h @@ -73,11 +73,8 @@ inline void PrefixSum(const int* src, int n, int* dst, cudaStream_t st) return BatchPrefixSum(&src, &n, &dst, 1, st); } -void AppendTokenIds(int** token_ids_ptrs, // - const int* output_ids, - const int* positions, - int batch_size, - cudaStream_t stream); +void invokeAppendOneTokenAndAdvanceSequence( + int* const* token_ids_ptrs, const int* selected_tokens, int* sequence_length, int batch_size, cudaStream_t stream); // Apply sigmoid gating: attn[i] *= sigmoid(gate[i]) // attn: [num_tokens, dim], contiguous diff --git a/src/turbomind/models/llama/llama_rope.h b/src/turbomind/models/llama/llama_rope.h index ec45bc103e..acbf71c407 100644 --- a/src/turbomind/models/llama/llama_rope.h +++ b/src/turbomind/models/llama/llama_rope.h @@ -39,9 +39,10 @@ struct Llama3RopeKernelParam { struct MropeRopeKernelParam { int3 section; + int stride{}; // per-batch row stride of the [batch, rows, 3] ids table (legacy per-batch layout) int* position_ids{}; int* position_delta{}; - int* position_offsets{}; + int* position_offsets{}; // per-token row offsets into a flat [rows, 3] table (no per-batch tables) int* length{}; }; diff --git a/src/turbomind/models/llama/llama_utils.h b/src/turbomind/models/llama/llama_utils.h index 75503673da..8ea7a054b1 100644 --- a/src/turbomind/models/llama/llama_utils.h +++ b/src/turbomind/models/llama/llama_utils.h @@ -1,7 +1,6 @@ // Copyright (c) OpenMMLab. All rights reserved. #pragma once -#include "src/turbomind/utils/nvtx_utils.h" #include #include #include @@ -65,18 +64,6 @@ size_t curandStateGetSize(); bool isDebug(); -struct NvtxScope { - explicit NvtxScope(const std::string& name) - { - PUSH_RANGE(name.c_str()); - } - - ~NvtxScope() - { - POP_RANGE; - } -}; - int64_t& gSequenceIds(int batch_idx); } // namespace turbomind diff --git a/src/turbomind/models/llama/unified_attention_layer.cc b/src/turbomind/models/llama/unified_attention_layer.cc index 8a466479bb..b2445240db 100644 --- a/src/turbomind/models/llama/unified_attention_layer.cc +++ b/src/turbomind/models/llama/unified_attention_layer.cc @@ -37,6 +37,7 @@ #include "src/turbomind/kernels/attention/attention.h" #include "src/turbomind/kernels/attention/decoding.h" #include "src/turbomind/kernels/attention/kv_cache_utils_v2.h" +#include "src/turbomind/kernels/attention/verification/attention.h" #include "src/turbomind/kernels/norm/rms_norm.h" #include "src/turbomind/macro.h" @@ -86,12 +87,12 @@ struct BlockConfig { struct AttentionData { struct Stat { - int n; - int q_sum; - int q_max; - int k_sum; - int k_max; - } decode, prefill; + int request_count; + int query_count; + int max_query_length; + int key_capacity_sum; + int max_key_capacity; + } verification, decode, prefill; Buffer_ block_ptrs; Buffer_ block_ptrs_offsets; @@ -141,7 +142,24 @@ UnifiedAttentionLayer::UnifiedAttentionLayer(std::vector weigh is_warm_up_{*context.is_warm_up}, context_{context}, linear_(*context.linear), - arch_{getSMVersion()} + arch_{getSMVersion()}, + sm_count_{getSMCount()}, + direct_verification_supported_{[&] { + const auto& reference = *weights.at(0); + const verification_attention::Capability capability{ + arch_, + reference.data_type, + reference.head_dim, + engine.spec_num_draft_tokens + 1, + quant_policy_, + engine.attn_cp_size, + reference.is_mla(), + static_cast(reference.sinks), + static_cast(context.device_prop.sharedMemPerBlockOptin), + }; + return engine.spec_num_draft_tokens > 0 + && verification_attention::supports(capability); + }()} { TM_CHECK_GE(weights.size(), 1); @@ -264,13 +282,29 @@ void UnifiedAttentionLayer::Run(BatchOp op, int phase, TensorMap& env) // Borrow the global mask owned by LanguageModel (pointer only; its content is // built at Forward time) and resolve this rank's token offset within it. - d->token_mask = env.at("token_mask").buffer().borrow(); - d->token_mask_base = 0; - if (engine_param_.attn_dp_size > 1) { - const auto& local_token_num = env.at("batch").data()[0]->local_token_num; - TM_CHECK_EQ((int)local_token_num.size(), engine_param_.attn_dp_size); - d->token_mask_base = - std::accumulate(local_token_num.begin(), local_token_num.begin() + engine_param_.attn_dp_rank, 0); + // Optional: the speculative executor composition does not produce one. + if (auto mask = env.try_("token_mask")) { + d->token_mask = mask->buffer().borrow(); + d->token_mask_base = 0; + if (engine_param_.attn_dp_size > 1) { + const auto& local_token_num = env.at("batch").data()[0]->local_token_num; + TM_CHECK_EQ((int)local_token_num.size(), engine_param_.attn_dp_size); + d->token_mask_base = + std::accumulate(local_token_num.begin(), local_token_num.begin() + engine_param_.attn_dp_rank, 0); + } + } + + // This is needed in async mode to clear the `attn` buffer for the finished sequences. Ohterwise random NaNs + // will crash the MoE router later + /// TODO: use better solution, this increase memory usage and heterogenous attention layers may still break it + if (tmp_attn_) { + Clear(tmp_attn_.slice( + 0, + d->verification.query_count + d->decode.query_count + d->prefill.query_count)); + Clear(split_cnt_); + if (engine_param_.attn_cp_size > 1) { + invokeFillNegInfML(partial_ML_.data(), partial_ML_.size() / 2, core::Context::stream().handle()); + } } } } @@ -304,31 +338,65 @@ void UnifiedAttentionLayer::Setup(int phase, TensorMap& env) copy(block_ptrs_offsets_buf_, bsz + 1, d.block_ptrs_offsets); } - /// prepare Q/K stats for decode/prefill - d.decode = d.prefill = {}; - - d.decode.n = std::find_if(rc.begin(), rc.end(), [](auto r) { return r->input_len > 1; }) - rc.begin(); - d.prefill.n = bsz - d.decode.n; + /// prepare Q/K stats for verification/decode/prefill + d.verification = d.decode = d.prefill = {}; + + if (direct_verification_supported_) { + d.verification.request_count = + std::find_if(rc.begin(), rc.end(), [](const Sequence* request) { + return !request->submitted->is_verification_row(); + }) + - rc.begin(); + d.decode.request_count = + std::find_if(rc.begin() + d.verification.request_count, + rc.end(), + [](const Sequence* request) { return request->submitted->input_len > 1; }) + - (rc.begin() + d.verification.request_count); + d.prefill.request_count = bsz - d.verification.request_count - d.decode.request_count; + } + else if (engine_param_.spec_num_draft_tokens > 0) { + d.prefill.request_count = bsz; + } + else { + d.decode.request_count = + std::find_if(rc.begin(), rc.end(), [](const Sequence* request) { + return request->submitted->input_len > 1; + }) + - rc.begin(); + d.prefill.request_count = bsz - d.decode.request_count; + } // d.dbg_offset = d.dbg_size = 0; for (int i = 0; i < bsz; ++i) { - const auto& c = *rc[i]; + const Sequence& c = *rc[i]; + const SubmittedRow& row = *c.submitted; // if (c.request->id == 4 && c.input_len > 1) { - // d.dbg_offset = d.decode.q_sum + d.prefill.q_sum; + // d.dbg_offset = d.decode.query_count + d.prefill.query_count; // d.dbg_size = c.input_len; // } - auto& s = i < d.decode.n ? d.decode : d.prefill; - s.q_sum += c.input_len; - s.k_sum += c.history_len + c.inflight_input_len + c.input_len; - s.q_max = std::max(s.q_max, c.input_len); - s.k_max = std::max(s.k_max, c.history_len + c.inflight_input_len + c.input_len); + auto& s = i < d.verification.request_count + ? d.verification + : (i < d.verification.request_count + d.decode.request_count + ? d.decode + : d.prefill); + s.query_count += row.input_len; + s.key_capacity_sum += row.key_capacity_end; + s.max_query_length = std::max(s.max_query_length, row.input_len); + s.max_key_capacity = std::max(s.max_key_capacity, row.key_capacity_end); } // auto &D = d.decode, &P = d.prefill; - // dbg(D.n, D.k_sum, D.k_max, P.n, P.q_sum, P.q_max, P.k_sum, P.k_max); + // dbg(D.request_count, + // D.key_capacity_sum, + // D.max_key_capacity, + // P.request_count, + // P.query_count, + // P.max_query_length, + // P.key_capacity_sum, + // P.max_key_capacity); /// handling different RoPE types if (rope_param_.type == RopeType::kDynamic) { @@ -362,6 +430,36 @@ void UnifiedAttentionLayer::Setup(int phase, TensorMap& env) } } +void UnifiedAttentionLayer::SetForwardMetadata(int phase, const AttentionForwardMetadata& metadata) +{ + auto& d = *data_[phase]; + + d.verification = { + metadata.verification.request_count, + metadata.verification.query_count, + metadata.verification.max_query_length, + metadata.verification.key_capacity_sum, + metadata.verification.max_key_capacity, + }; + d.decode = { + metadata.decode.request_count, + metadata.decode.query_count, + metadata.decode.max_query_length, + metadata.decode.key_capacity_sum, + metadata.decode.max_key_capacity, + }; + d.prefill = { + metadata.prefill.request_count, + metadata.prefill.query_count, + metadata.prefill.max_query_length, + metadata.prefill.key_capacity_sum, + metadata.prefill.max_key_capacity, + }; + + d.q_offsets = metadata.q_offsets.borrow(); + d.k_offsets = metadata.k_offsets.borrow(); +} + void UnifiedAttentionLayer::Forward(ForwardParam p) { TM_FUNCTION_SCOPE(); @@ -454,18 +552,30 @@ Tensor UnifiedAttentionLayer::core_attention(Tensor& qkv, const ForwardParam& p, auto& d = *data_.at(p.phase); - const int batch_size = d.decode.n + d.prefill.n; - const int q_count = qkv.shape(0); - - TM_CHECK_EQ(d.prefill.q_sum + d.decode.n, q_count); + const int query_count = + d.verification.query_count + d.decode.query_count + d.prefill.query_count; - const int local_q_kv_head_num = local_head_num + 2 * local_kv_head_num; + Tensor attn; + if (tmp_attn_) { + attn = tmp_attn_.slice(0, query_count); + } + else { + attn = {{query_count, local_head_num * size_per_head}, dtype, device}; + } - Tensor attn{{q_count, local_head_num * size_per_head}, dtype, device}; + if (query_count == 0) { + return attn; + } const bool is_mla = weights.is_mla(); - Tensor tmp_kv{{local_kv_head_num, is_mla ? 1 : 2, d.prefill.k_sum + MAX_CTA_S, size_per_head}, dtype, device}; + Tensor tmp_kv; + if (d.prefill.request_count) { + tmp_kv = Tensor{ + {local_kv_head_num, is_mla ? 1 : 2, d.prefill.key_capacity_sum + MAX_CTA_S, size_per_head}, + dtype, + device}; + } auto CreateParams = [&](int offset, AttentionData::Stat stat, int max_kv_splits, cudaStream_t stream) { AttentionParams params{}; @@ -497,11 +607,11 @@ Tensor UnifiedAttentionLayer::core_attention(Tensor& qkv, const ForwardParam& p, params.v_bias = params.k_bias + local_kv_head_num * size_per_head; } - params.batch_size = stat.n; + params.batch_size = stat.request_count; - params.token_num = stat.q_sum; - params.max_q_len = stat.q_max; - params.max_k_len = stat.k_max; + params.token_num = stat.query_count; + params.max_q_len = stat.max_query_length; + params.max_k_len = stat.max_key_capacity; TM_CHECK_LE(weights.cache_block_offset, INT_MAX); @@ -512,24 +622,27 @@ Tensor UnifiedAttentionLayer::core_attention(Tensor& qkv, const ForwardParam& p, engine_param_.cache_block_seq_len}; // prefill only - if (is_mla) { - params.linear_iter_params = LinearIteratorParams{ - tmp_kv.raw_data(), // flattened KV - stat.k_sum * size_per_head, // stride to next head - 0 // stride from K to V - }; - } - else { - params.linear_iter_params = LinearIteratorParams{ - tmp_kv.raw_data(), // flattened KV - stat.k_sum * size_per_head * 2, // stride to next head - stat.k_sum * size_per_head // stride from K to V - }; + if (tmp_kv) { + if (is_mla) { + params.linear_iter_params = LinearIteratorParams{ + tmp_kv.raw_data(), // flattened KV + stat.key_capacity_sum * size_per_head, // stride to next head + 0 // stride from K to V + }; + } + else { + params.linear_iter_params = LinearIteratorParams{ + tmp_kv.raw_data(), // flattened KV + stat.key_capacity_sum * size_per_head * 2, // stride to next head + stat.key_capacity_sum * size_per_head // stride from K to V + }; + } } - params.finished = d.finished.data() + offset; + const bool* finished = d.finished.data_or((bool*)nullptr); + params.finished = finished ? finished + offset : nullptr; // decode rows: base; prefill rows: + decode.n (this rank's slice of the global mask) - params.token_mask = d.token_mask.data() + d.token_mask_base + offset; + params.token_mask = d.token_mask ? d.token_mask.data() + d.token_mask_base + offset : nullptr; params.cu_q_len = d.q_offsets.data() + offset; params.cu_k_len = d.k_offsets.data() + offset; params.readonly_block_num = d.readonly_block_num.data() + offset; @@ -608,19 +721,69 @@ Tensor UnifiedAttentionLayer::core_attention(Tensor& qkv, const ForwardParam& p, return params; }; + auto MakeVerificationArguments = [&](const AttentionParams& writer, + int query_offset, + AttentionData::Stat stat, + cudaStream_t stream) { + verification_attention::Arguments a{}; + a.out = writer.out; + a.q = writer.q; + a.q_bias = writer.q_bias; + a.q_stride = writer.stride; + a.block_ptrs = writer.block_iter_params.block_ptrs; + a.block_ptr_offsets = writer.block_iter_params.cu_block_nums; + a.q_offsets = writer.cu_q_len; + a.k_offsets = writer.cu_k_len; + a.finished = writer.finished; + a.request_count = stat.request_count; + a.query_count = stat.query_count; + a.query_offset = query_offset; + a.max_query_length = stat.max_query_length; + a.max_key_length = stat.max_key_capacity; + a.query_head_count = writer.num_heads; + a.kv_head_count = writer.num_kv_heads; + a.query_group_size = a.query_head_count / a.kv_head_count; + a.query_group_size_divmod = cutlass::FastDivmod(a.query_group_size); + a.head_dim = writer.size_per_head; + a.block_len = writer.block_iter_params.block_len; + a.block_len_divmod = cutlass::FastDivmod(a.block_len); + a.cache_block_offset = writer.block_iter_params.offset; + a.window_size = writer.window_size; + a.qk_scale_log2 = writer.inv_sqrt_dh; + a.rope = writer.rope_param; + a.partial_o = partial_O_.data(); + a.partial_ml = partial_ML_.data(); + a.data_type = dtype; + a.stream = stream; + + const int m_slices = cdiv( + a.max_query_length * a.query_group_size, + verification_attention::CtaM(a)); + const int base_ctas = a.request_count * a.kv_head_count * m_slices; + a.split_count = verification_attention::choose_split_count(a.query_count, + base_ctas, + a.max_key_length, + verification_attention::KeyTile(a), + kMaxWorkspaceTokens, + kMaxKVSplits, + sm_count_); + return a; + }; + const cudaStream_t stream = core::Context::stream().handle(); cudaStream_t pf_stream = stream; cudaStream_t dc_stream = stream; - if (d.decode.n && d.prefill.n) { + const bool has_executor_attention = d.verification.request_count || d.decode.request_count; + if (has_executor_attention && d.prefill.request_count) { pf_stream = aux_stream_; TM_CUDA_CHECK(cudaEventRecord(qkv_event_, stream)); TM_CUDA_CHECK(cudaStreamWaitEvent(aux_stream_, qkv_event_)); } - if (d.prefill.n && !is_warm_up_) { - const int offset = d.decode.n; + if (d.prefill.request_count && !is_warm_up_) { + const int offset = d.verification.request_count + d.decode.request_count; // We are executing prefill & decoding kernels concurrently, but only have 1 workspace // disable split kv for prefill for now auto params = CreateParams(offset, d.prefill, 1, pf_stream); @@ -629,7 +792,7 @@ Tensor UnifiedAttentionLayer::core_attention(Tensor& qkv, const ForwardParam& p, TM_CUDA_CHECK(cudaGetLastError()); /// TODO: skip flattening for `sm_80` - invokeFlattenKV_v2_(params, d.prefill.k_sum); + invokeFlattenKV_v2_(params, d.prefill.key_capacity_sum); TM_CUDA_CHECK(cudaGetLastError()); dispatchAttention(params); @@ -637,15 +800,28 @@ Tensor UnifiedAttentionLayer::core_attention(Tensor& qkv, const ForwardParam& p, } } - if (d.decode.n && !is_warm_up_) { - auto params = CreateParams(0, d.decode, kMaxKVSplits, dc_stream); + if (d.verification.request_count && !is_warm_up_) { + auto params = CreateParams(0, d.verification, kMaxKVSplits, dc_stream); + if constexpr (sizeof(T) == 2) { + invokeProcessKV_v2_(params); + TM_CUDA_CHECK(cudaGetLastError()); + auto arguments = MakeVerificationArguments(params, 0, d.verification, dc_stream); + verification_attention::run(arguments); + TM_CUDA_CHECK(cudaGetLastError()); + } + } + + if (d.decode.request_count && !is_warm_up_) { + const int offset = d.verification.request_count; + auto params = CreateParams( + offset, d.decode, d.verification.request_count ? 1 : kMaxKVSplits, dc_stream); if constexpr (sizeof(T) == 2) { dispatchDecoding(params); TM_CUDA_CHECK(cudaGetLastError()); } } - if (d.decode.n && d.prefill.n) { + if (has_executor_attention && d.prefill.request_count) { TM_CUDA_CHECK(cudaEventRecord(aux_event_, aux_stream_)); TM_CUDA_CHECK(cudaStreamWaitEvent(stream, aux_event_)); } diff --git a/src/turbomind/models/llama/unified_attention_layer.h b/src/turbomind/models/llama/unified_attention_layer.h index 3e0eb77a62..6a052cfe5b 100644 --- a/src/turbomind/models/llama/unified_attention_layer.h +++ b/src/turbomind/models/llama/unified_attention_layer.h @@ -39,6 +39,23 @@ namespace turbomind { struct AttentionData; +struct AttentionForwardMetadata { + struct Partition { + int request_count; + int query_count; + int max_query_length; + int key_capacity_sum; + int max_key_capacity; + }; + + Partition verification; + Partition decode; + Partition prefill; + + Buffer_ q_offsets; + Buffer_ k_offsets; +}; + class UnifiedAttentionLayer { public: using WeightType = AttentionWeight; @@ -66,6 +83,8 @@ class UnifiedAttentionLayer { void Forward(ForwardParam p); + void SetForwardMetadata(int phase, const AttentionForwardMetadata& metadata); + private: void Setup(int phase, TensorMap& env); @@ -86,6 +105,8 @@ class UnifiedAttentionLayer { LlamaLinear& linear_; const int arch_{}; + const int sm_count_{}; + const bool direct_verification_supported_{}; cudaStream_t aux_stream_; cudaEvent_t qkv_event_; @@ -107,6 +128,7 @@ class UnifiedAttentionLayer { Tensor_ partial_O_; Tensor_ partial_ML_; Tensor_ split_cnt_; + Tensor tmp_attn_; Buffer_ rope_base_buf_; Buffer_ mrope_default_buf_; diff --git a/src/turbomind/models/llama/unified_decoder.cc b/src/turbomind/models/llama/unified_decoder.cc index 3eee72c496..2602b577be 100644 --- a/src/turbomind/models/llama/unified_decoder.cc +++ b/src/turbomind/models/llama/unified_decoder.cc @@ -17,6 +17,7 @@ #include "src/turbomind/models/llama/unified_attention_layer.h" #include "src/turbomind/models/llama/unified_decoder.h" #include "src/turbomind/models/model_weight.h" +#include "src/turbomind/models/speculative/hidden_state_tap.h" #include "src/turbomind/utils/anomaly_handler.h" #include "src/turbomind/utils/cuda_utils.h" @@ -180,7 +181,10 @@ void UnifiedDecoder::AllreduceResidualRMSnorm(Tensor& hidden_states, } } -void UnifiedDecoder::Forward(int phase, TensorMap& args, const std::vector& weights) +void UnifiedDecoder::Forward(int phase, + TensorMap& args, + const std::vector& weights, + const Tensor& selected_hidden_buffer) { TM_FUNCTION_SCOPE(); /** @@ -203,11 +207,38 @@ void UnifiedDecoder::Forward(int phase, TensorMap& args, const std::vector()[0]->local_token_num; + HiddenStateTap* tap = nullptr; + Tensor handle = args.try_consume("hidden_state_tap"); + if (handle) { + tap = handle.data()[0]; + } + + Tensor local_residual = args.try_consume("residual"); + + std::vector token_topology; + if (const Tensor* topology = args.try_("decoder_local_token_nums")) { + token_topology.assign(topology->data(), topology->data() + topology->shape(0)); + } + else { + token_topology = args.at("batch").data()[0]->local_token_num; + } + const int* local_token_nums = token_topology.data(); const auto local_token_num = local_residual.shape(0); - const auto global_token_num = std::accumulate(local_token_nums.begin(), local_token_nums.end(), ssize_t{}); + const auto global_token_num = std::accumulate(token_topology.begin(), token_topology.end(), ssize_t{}); + + // The MoE router consumes a per-token validity mask (attn-DP padding rows route + // nowhere). This executor composition has no padding rows — supply an all-valid + // mask when no producer provided one. + const bool* token_mask = nullptr; + if (const Tensor* mask = args.try_("token_mask")) { + token_mask = (const bool*)mask->buffer().raw_data(); + } + else if (global_token_num > 0) { + all_valid_mask_ = Buffer_{(size_t)global_token_num, kDEVICE}; + TM_CUDA_CHECK(cudaMemsetAsync(all_valid_mask_.data(), 1, global_token_num, core::Context::stream().handle())); + token_mask = all_valid_mask_.data(); + } TM_CHECK_EQ(local_token_num, local_token_nums[attn_dp_rank_]); @@ -224,9 +255,8 @@ void UnifiedDecoder::Forward(int phase, TensorMap& args, const std::vector 1) { // Offset hidden states buffer for mixed DP - TM_CHECK_EQ(local_token_nums.size(), attn_dp_size_); std::vector offsets(attn_dp_size_ + 1, 0); - std::inclusive_scan(local_token_nums.data(), local_token_nums.data() + attn_dp_size_, offsets.begin() + 1); + std::inclusive_scan(local_token_nums, local_token_nums + attn_dp_size_, offsets.begin() + 1); const int offset = offsets[attn_dp_rank_]; local_hidden_states = global_hidden_states.slice({offset, 0}, {local_token_num, -1}); @@ -242,17 +272,23 @@ void UnifiedDecoder::Forward(int phase, TensorMap& args, const std::vectorattention_norm; - invokeRMSNorm(local_hidden_states, - local_residual, - first_norm.weight, - first_norm.norm_eps_, - first_norm.zero_centered_, - stream); - - TM_CUDA_CHECK(cudaGetLastError()); + Tensor layer_attention_input; + if (args.contains("attention_input")) { + layer_attention_input = args.try_consume("attention_input"); + } + else { + const auto& first_norm = *weights.at(0)->attention_norm; + invokeRMSNorm(local_hidden_states, + local_residual, + first_norm.weight, + first_norm.norm_eps_, + first_norm.zero_centered_, + stream); + + layer_attention_input = local_hidden_states; + } - TM_DEBUG_TENSOR(local_hidden_states, Concat("norm0", 0), 2); + TM_DEBUG_TENSOR(layer_attention_input, Concat("norm0", 0), 2); // auto stack_alloc{core::Context::device_alloc().adapt()}; // core::ContextGuard ctx{Allocator{stack_alloc}}; @@ -276,11 +312,11 @@ void UnifiedDecoder::Forward(int phase, TensorMap& args, const std::vectorlinear_attn) { linear_attn_layer_->Forward( - {phase, local_hidden_states, local_hidden_states, weights.at(layer)->linear_attn.get()}); + {phase, layer_attention_input, local_hidden_states, weights.at(layer)->linear_attn.get()}); } else { auto* attn = weights.at(layer)->attention.get(); - attn_layer_->Forward({phase, local_hidden_states, local_hidden_states, attn, layer}); + attn_layer_->Forward({phase, layer_attention_input, local_hidden_states, attn, layer}); } TM_DEBUG_TENSOR(local_hidden_states, Concat("attn_block", layer), 2); @@ -308,8 +344,8 @@ void UnifiedDecoder::Forward(int phase, TensorMap& args, const std::vectormoe_ffn) { moe_ffn_layer_->Forward({global_hidden_states, global_hidden_states, - local_token_nums, + token_topology, weights.at(layer)->moe_ffn.get(), (int)layer, - (const bool*)args.at("token_mask").buffer().raw_data()}); + token_mask}); } if (ffn_layer_ && weights.at(layer)->feed_forward) { @@ -354,13 +390,24 @@ void UnifiedDecoder::Forward(int phase, TensorMap& args, const std::vectorTapOrdinal(completed_layer_count); + if (tap_ordinal >= 0) { + tap->Capture( + tap_ordinal, local_residual, local_hidden_states, core::Context::stream().handle()); + } + } + TM_DEBUG_TENSOR(local_residual, Concat("residual1", layer), 2); TM_DEBUG_TENSOR(local_hidden_states, Concat("norm0", layer + 1), 2); + layer_attention_input = local_hidden_states; + // if (layer == layer_num_ - 1) { // args.at("batch").data()[0]->Notify(); // } @@ -372,11 +419,12 @@ void UnifiedDecoder::Forward(int phase, TensorMap& args, const std::vector(selected_hidden_buffer); + const bool output_hidden_states = args.try_("output_hidden_states"); Tensor hidden_states{local_hidden_states}; - if (d_comm_ && (output_hidden_states || reuse_hidden_states)) { + if (!caller_owns_selected_states && d_comm_ && (output_hidden_states || reuse_hidden_states)) { // The full `hidden_states` buffer is needed for output but it's a ref into `symm_buf` atm. // Copy to residual buf so that `symm_buf` may be reused safely later Copy(hidden_states, local_residual); @@ -384,15 +432,25 @@ void UnifiedDecoder::Forward(int phase, TensorMap& args, const std::vector& weights); + // `selected_hidden_buffer` is the caller's request to write selected + // hidden states into its own buffer (and receive `pre_final_residual`); + // empty leaves selection to the decoder. The typed argument is the + // request — env keys carry no activation. + void Forward(int phase, + TensorMap& env, + const std::vector& weights, + const Tensor& selected_hidden_buffer = {}); + + void CommitAcceptedState(int phase, const Buffer_& accept_len) + { + if (linear_attn_layer_) { + linear_attn_layer_->CommitAcceptedState(phase, accept_len); + } + } + + size_t SpeculativeStateJournalBytes(int request_count, int verification_positions) const + { + return linear_attn_layer_ ? + linear_attn_layer_->SpeculativeStateJournalBytes(request_count, verification_positions) : + 0; + } + + void SetAttentionForwardMetadata(int phase, const AttentionForwardMetadata& metadata) + { + attn_layer_->SetForwardMetadata(phase, metadata); + } private: const size_t layer_num_; @@ -45,6 +71,9 @@ class UnifiedDecoder { // Per-layer post-FFN reduce group, precomputed in the constructor. std::vector ffn_group_; + // All-valid per-token mask, materialized when no producer supplies `token_mask`. + Buffer_ all_valid_mask_; + comm::DeviceCommImpl* const d_comm_; const int tune_layer_num_; diff --git a/src/turbomind/models/model_root.h b/src/turbomind/models/model_root.h index 31769368bf..92e03740d6 100644 --- a/src/turbomind/models/model_root.h +++ b/src/turbomind/models/model_root.h @@ -45,6 +45,11 @@ class ModelRoot: public core::Module { return text_model.get(); } + ModelWeight* draft_model_ptr() const + { + return draft_model.get(); + } + /// Convenience accessor for the optional VLM sub-tree. Nullptr for /// text-only checkpoints (the spec never attached a vision root). VisionModelWeight* vision_model_ptr() const @@ -54,6 +59,7 @@ class ModelRoot: public core::Module { #define MODEL_ROOT_CHILDREN(X) \ X(ModelWeight, text_model) \ + X(ModelWeight, draft_model) \ X(VisionModelWeight, vision_model) #define MODEL_ROOT_PARAMS(X) diff --git a/src/turbomind/models/model_weight.cc b/src/turbomind/models/model_weight.cc index a359b526b6..0621e0990b 100644 --- a/src/turbomind/models/model_weight.cc +++ b/src/turbomind/models/model_weight.cc @@ -8,7 +8,11 @@ namespace turbomind { ModelWeight::ModelWeight(const core::ModelWeightConfig& cfg): - tp_size(cfg.tp_size), tp_rank(cfg.tp_rank), data_type(cfg.data_type), hidden_units(cfg.hidden_units) + data_type(cfg.data_type), + hidden_units(cfg.hidden_units), + tp_size(cfg.tp_size), + tp_rank(cfg.tp_rank), + decoder_only(cfg.decoder_only) { } @@ -39,17 +43,21 @@ void ModelWeight::prepare() head_dim = attn_layer->attention->head_dim; kv_head_num = attn_layer->attention->kv_head_num; - vocab_size = tok_embeddings.shape(0); - embedding_size = vocab_size; - num_layer = layers->size(); - vocab_size_padded = TM_CHECK_NOTNULL(output)->output_dim * tp_size; + num_layer = layers->size(); + if (!decoder_only) { + vocab_size = tok_embeddings.shape(0); + embedding_size = vocab_size; + vocab_size_padded = TM_CHECK_NOTNULL(output)->output_dim * tp_size; + } layer_types.resize(num_layer); for (int i = 0; i < num_layer; ++i) { layer_types[i] = layer(i)->linear_attn ? 1 : 0; } - EnsureFloatDtype(tok_embeddings, data_type); + if (!decoder_only) { + EnsureFloatDtype(tok_embeddings, data_type); + } } DecoderLayerWeight* ModelWeight::layer(int i) const @@ -78,7 +86,7 @@ std::vector ModelWeight::layers_list() const bool ModelWeight::verify(std::vector& missing) { Module::verify(missing); - if (!tok_embeddings) { + if (!decoder_only && !tok_embeddings) { missing.push_back(full_path() + ": missing tok_embeddings"); } if (!norm) { diff --git a/src/turbomind/models/model_weight.h b/src/turbomind/models/model_weight.h index 1612cad30b..b433cfac98 100644 --- a/src/turbomind/models/model_weight.h +++ b/src/turbomind/models/model_weight.h @@ -18,7 +18,8 @@ struct ModelWeightConfig: ModuleConfig { X(int, tp_size) \ X(int, tp_rank) \ X(DataType, data_type) \ - X(int, hidden_units) + X(int, hidden_units) \ + X(bool, decoder_only, false) MODEL_WEIGHT_FIELDS(TM_MEMBER) TM_FOR_EACH(ModelWeightConfig, MODEL_WEIGHT_FIELDS) @@ -51,6 +52,7 @@ class ModelWeight: public core::Module { #define MODEL_WEIGHT_CHILDREN(X) \ X(LinearWeight, output) \ X(NormWeight, norm) \ + X(core::Module, spec) \ X(core::ModuleList, layers) \ X(core::ModuleList, meta_experts) @@ -76,6 +78,7 @@ class ModelWeight: public core::Module { // --- From ModelWeightConfig at construction --- int tp_size{}; int tp_rank{}; + bool decoder_only{}; private: mutable std::vector layers_cache_; diff --git a/src/turbomind/models/output_processor.cc b/src/turbomind/models/output_processor.cc index 577b152ec6..2e04e10729 100644 --- a/src/turbomind/models/output_processor.cc +++ b/src/turbomind/models/output_processor.cc @@ -1,10 +1,10 @@ #include "src/turbomind/models/output_processor.h" -#include - #include "src/turbomind/engine/request.h" #include "src/turbomind/kernels/cross_entropy_kernels.h" +#include "src/turbomind/models/language_model.h" +#include "src/turbomind/models/model_weight.h" // #include "dbg.h" @@ -18,14 +18,12 @@ struct OutputProcessor::Impl { static constexpr auto kAll = GenerationConfig::kAll; - const int vocab_size_; - const int max_logits_len_; - const int tp_rank_; - - std::function lm_head_; + LanguageModel& model_; + const int vocab_size_; + const int tp_rank_; - Impl(int vocab_size, int max_logits_len, int tp_rank, int phases, std::function lm_head): - vocab_size_{vocab_size}, max_logits_len_{max_logits_len}, tp_rank_{tp_rank}, lm_head_{std::move(lm_head)} + Impl(LanguageModel& model, int tp_rank, int phases): + model_{model}, vocab_size_{model.weights().vocab_size}, tp_rank_{tp_rank} { for (int i = 0; i < phases; ++i) { data_.emplace_back(); @@ -123,11 +121,12 @@ struct OutputProcessor::Impl { vector sel_tokens; bool has_ce = false; for (int i = 0; i < rc.size(); ++i) { - using Size = Interval::Size; - auto& c = *rc[i]; - all_tokens.emplace_back(c.history_len + c.inflight_input_len, Size{c.input_len}); - sel_tokens.emplace_back(c.history_len + c.inflight_input_len + c.input_len - 1, Size{1}); - if (!c.generating) { + using Size = Interval::Size; + auto& c = *rc[i]; + const SubmittedRow& row = *c.submitted; + all_tokens.emplace_back(row.history_len + c.inflight_input_len, Size{row.input_len}); + sel_tokens.emplace_back(row.history_len + c.inflight_input_len + row.input_len - 1, Size{1}); + if (!row.generating) { sel_tokens.back() = {}; } has_ce = has_ce || (bool)c.input_ce_loss; @@ -158,8 +157,9 @@ struct OutputProcessor::Impl { int offset = 0; for (int i = 0; i < rc.size(); ++i) { - auto& c = *rc[i]; - auto& g = c.req->gen_cfg; + auto& c = *rc[i]; + const SubmittedRow& row = *c.submitted; + auto& g = c.req->gen_cfg; if (c.output_hidden_states) { Matching m{c.output_hidden_states, c.hidden_states_offset}; int type = 0; @@ -205,7 +205,7 @@ struct OutputProcessor::Impl { d.ce_loss_segments.push_back({c.req, c.ce_loss, m.src, !c.input_ce_loss}); } } - offset += c.input_len; + offset += row.input_len; } // logits depends on hidden states @@ -271,9 +271,9 @@ struct OutputProcessor::Impl { } } - void ComputeAndOutputLogits(Data& data, const Tensor& h) + void ComputeAndOutputLogits(Data& data, const Tensor& h, const TensorMap& env) { - const int step_size = max_logits_len_; + const int step_size = model_.max_logits_len(env); // Coroutine frame int p = 0; @@ -289,7 +289,7 @@ struct OutputProcessor::Impl { if (auto chunk = r & Interval{r.begin(), Size{step_size}}) { // dbg(&chunk); // Compute full logits by chunks - auto logits = lm_head_(h.slice(chunk.begin(), (int)chunk.size())); + auto logits = model_.Logits(h.slice(chunk.begin(), (int)chunk.size()), {}, env); if (!success) { success = OutputLogitsImpl(ranges, p, logits, chunk.begin(), 2); } @@ -364,7 +364,7 @@ struct OutputProcessor::Impl { OutputHiddenStates(d.output_states, hidden_states, 2); } if (d.full_logits || d.full_ce_loss) { - ComputeAndOutputLogits(d, hidden_states); + ComputeAndOutputLogits(d, hidden_states, env); } } @@ -381,9 +381,8 @@ struct OutputProcessor::Impl { OutputProcessor::~OutputProcessor() = default; -OutputProcessor::OutputProcessor( - int vocab_size, int max_logits_len, int tp_rank, int phases, std::function lm_head): - impl_{std::make_unique(vocab_size, max_logits_len, tp_rank, phases, std::move(lm_head))} +OutputProcessor::OutputProcessor(LanguageModel& model, int tp_rank, int phases): + impl_{std::make_unique(model, tp_rank, phases)} { } diff --git a/src/turbomind/models/output_processor.h b/src/turbomind/models/output_processor.h index 2dcd569d4c..dc455489b8 100644 --- a/src/turbomind/models/output_processor.h +++ b/src/turbomind/models/output_processor.h @@ -4,12 +4,13 @@ namespace turbomind { +class LanguageModel; + class OutputProcessor { public: ~OutputProcessor(); - OutputProcessor( - int vocab_size, int max_logits_len, int tp_rank, int phases, std::function lm_head); + OutputProcessor(LanguageModel& model, int tp_rank, int phases); void Run(BatchOp op, int phase, TensorMap& env); diff --git a/src/turbomind/models/qwenvit/qwenvit.cc b/src/turbomind/models/qwenvit/qwenvit.cc index 2f34673c5b..66b86b36f5 100644 --- a/src/turbomind/models/qwenvit/qwenvit.cc +++ b/src/turbomind/models/qwenvit/qwenvit.cc @@ -50,6 +50,7 @@ struct QwenVit::Impl { comm::DeviceCommImpl* const d_comm_; const int tp_group_; const DataType engine_data_type_; + const bool successor_embeddings_; const std::string communicator_; Buffer_ grid_thws_buf_; // (t, h, w) @@ -62,8 +63,8 @@ struct QwenVit::Impl { Tensor batch_input; int batch_size; std::vector> grid_thws_host; - std::vector> image_embeds_coords; // (size, pos) for image embeddings - std::vector> input_embeds_coords; // (size, pos) for input embeddings + std::vector target_patches; // (rows, source, destination) for target embeddings + std::vector successor_patches; // (rows, source, destination) for successor embeddings // for RoPE / pos-embed interpolation Tensor_ grid_thws; @@ -107,8 +108,8 @@ struct QwenVit::Impl { window_attn_batch_size = 0; max_window_attn_len = 0; grid_thws_host.clear(); - image_embeds_coords.clear(); - input_embeds_coords.clear(); + target_patches.clear(); + successor_patches.clear(); } }; @@ -117,7 +118,11 @@ struct QwenVit::Impl { std::vector data_; - Impl(const EngineParam& engine, const Context& ctx, const QwenVitWeight& weights, int phases): + Impl(const EngineParam& engine, + const Context& ctx, + const QwenVitWeight& weights, + int phases, + bool successor_embeddings): weights_{weights}, config_{weights.config()}, linear_{*ctx.linear}, @@ -125,6 +130,7 @@ struct QwenVit::Impl { d_comm_{ctx.comm.d_comm}, tp_group_{ctx.comm.d_tp_group}, engine_data_type_{engine.data_type}, + successor_embeddings_{successor_embeddings}, communicator_{engine.communicator} { for (int i = 0; i < phases; ++i) { @@ -193,22 +199,35 @@ struct QwenVit::Impl { int input_ids_offsets = 0; int image_embeds_offsets = 0; for (int i = 0; i < rc.size(); ++i) { - const auto& s = *rc[i]; + const Sequence& s = *rc[i]; + const SubmittedRow& submitted = *s.submitted; - if ((not s.autoregres) && (not s.multimodal_inputs.empty())) { + if ((not submitted.autoregres) && (not s.multimodal_inputs.empty())) { ++mm_prefill_seqs; images_total += (int)s.multimodal_inputs.size(); - Interval text{s.history_len + s.inflight_input_len, Interval::Size{s.input_len}}; + const int begin = submitted.history_len + s.inflight_input_len; + const int end = begin + submitted.input_len; + const Interval target{begin, end}; + const Interval successor{begin + 1, std::min(end + 1, s.seq_len)}; for (const auto& mm : s.multimodal_inputs) { - auto o = mm->interval & text; - if (auto size = (int)o.size()) { + const Interval target_overlap = mm->interval & target; + const Interval successor_overlap = successor_embeddings_ ? mm->interval & successor : Interval{}; + if (!target_overlap.empty() || !successor_overlap.empty()) { pixel_values.push_back(mm->data); d.batch_size += mm->data.shape(0); - const int text_offset = input_ids_offsets + o.begin() - text.begin(); - const int image_offset = image_embeds_offsets + o.begin() - mm->interval.begin(); - d.input_embeds_coords.emplace_back(size, text_offset); - d.image_embeds_coords.emplace_back(size, image_offset); + if (!target_overlap.empty()) { + d.target_patches.push_back( + {static_cast(target_overlap.size()), + image_embeds_offsets + target_overlap.begin() - mm->interval.begin(), + input_ids_offsets + target_overlap.begin() - target.begin()}); + } + if (!successor_overlap.empty()) { + d.successor_patches.push_back( + {static_cast(successor_overlap.size()), + image_embeds_offsets + successor_overlap.begin() - mm->interval.begin(), + input_ids_offsets + successor_overlap.begin() - successor.begin()}); + } const auto& grid_thw = mm->grid_thw; d.grid_thws_host.emplace_back(grid_thw); @@ -219,7 +238,7 @@ struct QwenVit::Impl { } } - input_ids_offsets += s.autoregres ? 1 : s.input_len; + input_ids_offsets += submitted.input_len; } // Prefix-cache observability: on a fully-cached image, the window filter @@ -612,11 +631,12 @@ struct QwenVit::Impl { int upper_segs = 0; int total_q_tokens = 0; for (int i = 0; i < bsz; ++i) { - const auto& s = *rc[i]; + const Sequence& s = *rc[i]; + const SubmittedRow& submitted = *s.submitted; d.mrope_offsets_host.data()[i] = total_q_tokens; - total_q_tokens += s.autoregres ? 1 : s.input_len; - if (!s.autoregres && !s.multimodal_inputs.empty()) { + total_q_tokens += submitted.input_len; + if (!submitted.autoregres && !s.multimodal_inputs.empty()) { upper_segs += 2 * (int)s.multimodal_inputs.size() + 1; } } @@ -635,11 +655,12 @@ struct QwenVit::Impl { int max_seg_len = 0; for (int i = 0; i < bsz; ++i) { - const auto& s = *rc[i]; - const int seq_len = (int)s.req->inputs.at("input_ids").shape(0); - const bool needs_table = !s.autoregres && !s.multimodal_inputs.empty(); - const int active_start = s.history_len + s.inflight_input_len; - const int active_end = active_start + s.input_len; + const Sequence& s = *rc[i]; + const SubmittedRow& submitted = *s.submitted; + const int seq_len = (int)s.req->inputs.at("input_ids").shape(0); + const bool needs_table = !submitted.autoregres && !s.multimodal_inputs.empty(); + const int active_start = submitted.history_len + s.inflight_input_len; + const int active_end = active_start + submitted.input_len; const int q_offset = d.mrope_offsets_host.data()[i]; auto emit = [&](int run_start, int run_n, int run_base, int h2, int w2) { @@ -682,7 +703,7 @@ struct QwenVit::Impl { emit(row, active_end - row, pos, /*h2=*/0, /*w2=*/0); } - d.mrope_length_host.data()[i] = needs_table ? s.input_len : 0; + d.mrope_length_host.data()[i] = needs_table ? submitted.input_len : 0; d.mrope_delta_host.data()[i] = mm_off; } @@ -884,8 +905,7 @@ struct QwenVit::Impl { // must match before publishing. EnsureFloatDtype(image_embeds, engine_data_type_); - args.produce("multimodal", - MultiModalEmbeddingData{image_embeds, d.image_embeds_coords, d.input_embeds_coords}.buf()); + args.produce("multimodal", MultiModalEmbeddingData{image_embeds, d.target_patches, d.successor_patches}.buf()); } template @@ -1044,8 +1064,12 @@ struct QwenVit::Impl { } }; -QwenVit::QwenVit(const EngineParam& engine, const Context& ctx, const QwenVitWeight& weights, int phases): - impl_{std::make_unique(engine, ctx, weights, phases)} +QwenVit::QwenVit(const EngineParam& engine, + const Context& ctx, + const QwenVitWeight& weights, + int phases, + bool successor_embeddings): + impl_{std::make_unique(engine, ctx, weights, phases, successor_embeddings)} { } diff --git a/src/turbomind/models/qwenvit/qwenvit.h b/src/turbomind/models/qwenvit/qwenvit.h index b69bf2dae9..84a84eb93a 100644 --- a/src/turbomind/models/qwenvit/qwenvit.h +++ b/src/turbomind/models/qwenvit/qwenvit.h @@ -24,7 +24,11 @@ class QwenVitWeight; /// - RMSNorm vs LayerNorm (Qwen2): norm_type class QwenVit: public VisionModel { public: - QwenVit(const EngineParam& engine, const Context& ctx, const QwenVitWeight& weights, int phases); + QwenVit(const EngineParam& engine, + const Context& ctx, + const QwenVitWeight& weights, + int phases, + bool successor_embeddings); ~QwenVit() override; diff --git a/src/turbomind/models/speculative/collect_hidden_states.cc b/src/turbomind/models/speculative/collect_hidden_states.cc new file mode 100644 index 0000000000..7263a3126f --- /dev/null +++ b/src/turbomind/models/speculative/collect_hidden_states.cc @@ -0,0 +1,82 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#include "src/turbomind/models/speculative/collect_hidden_states.h" + +#include "src/turbomind/comm/device_comm.h" +#include "src/turbomind/comm/padded_row_allgather.h" +#include "src/turbomind/core/check.h" +#include "src/turbomind/kernels/core/math.h" +#include "src/turbomind/models/llama/context.h" +#include "src/turbomind/models/llama/llama_params.h" + +namespace turbomind { + +CollectHiddenStates::CollectHiddenStates(const EngineParam& engine, + const Context& context, + int phases, + int capture_width, + int hidden_units, + DataType data_type): + capture_width_{capture_width}, + hidden_units_{hidden_units}, + capacity_{cdiv(engine.max_forward_token_num, engine.attn_cp_size * engine.attn_tp_size)}, + attn_dp_rank_{engine.attn_dp_rank}, + model_tp_group_{context.comm.d_tp_group}, + model_tp_rank_{engine.model_tp_rank}, + model_tp_size_{engine.attn_cp_size * engine.attn_tp_size}, + d_comm_{context.comm.d_comm}, + data_(phases) +{ + Allocator symmetric_allocator; + if (model_tp_size_ > 1) { + symmetric_allocator = GetSymmAllocator(context.comm.d_comm); + } + for (Data& data : data_) { + data.captured = Tensor{{capacity_, capture_width_}, data_type, kDEVICE}; + if (model_tp_size_ > 1) { + data.gathered_padded = { + {model_tp_size_ * capacity_, hidden_units_}, data_type, symmetric_allocator}; + } + } +} + +void CollectHiddenStates::Begin(int phase, const std::vector& local_token_nums) +{ + Data& data = data_.at(phase); + + const int global_rank = d_comm_ ? d_comm_->rank(0) : 0; + const int global_size = d_comm_ ? d_comm_->n_ranks(0) : 1; + const int model_tp_size = d_comm_ ? d_comm_->n_ranks(model_tp_group_) : 1; + data.owned = comm::ComputeTokenOwnership(global_rank, global_size, model_tp_size, local_token_nums.data()); + data.token_num = local_token_nums[attn_dp_rank_]; +} + +void CollectHiddenStates::SeedWarmup(int phase, cudaStream_t stream) +{ + Data& data = data_.at(phase); + const int rows = data.owned.row_count(); + if (rows > 0) { + TM_CUDA_CHECK(cudaMemsetAsync( + data.captured.raw_data(), 0, byte_size(data.captured.dtype(), size_t(rows) * capture_width_), stream)); + } +} + +Tensor CollectHiddenStates::Gather(int phase, const Tensor& local_buffer, cudaStream_t stream) +{ + Data& data = data_.at(phase); + const int rows = data.owned.row_count(); + if (model_tp_size_ == 1) { + return local_buffer.slice({0, 0}, {rows, local_buffer.shape(1)}); + } + return PaddedRowAllGather(local_buffer, + data.gathered_padded, + data.token_num, + rows, + model_tp_rank_, + model_tp_size_, + *d_comm_, + model_tp_group_, + stream); +} + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/collect_hidden_states.h b/src/turbomind/models/speculative/collect_hidden_states.h new file mode 100644 index 0000000000..b1d65bc6df --- /dev/null +++ b/src/turbomind/models/speculative/collect_hidden_states.h @@ -0,0 +1,73 @@ +// Copyright (c) OpenMMLab. All rights reserved. +#pragma once + +#include "src/turbomind/comm/device_comm.h" +#include "src/turbomind/comm/token_ownership.h" +#include "src/turbomind/core/core.h" + +#include + +namespace turbomind { + +class Context; +class EngineParam; + +/// Shared machinery of the hidden-state taps: the token rows this rank owns +/// during the target pass, the buffer that captures them, and the padded +/// all-gather that reunites them across model-TP ranks for the draft pass. +/// `capture_width` is the per-row width a tap writes (H, or L * H when several +/// target layers are tapped); `hidden_units` is the width the gather returns. +class CollectHiddenStates { +public: + CollectHiddenStates(const EngineParam& engine, + const Context& context, + int phases, + int capture_width, + int hidden_units, + DataType data_type); + + /// Records the token rows this rank owns for the round. + void Begin(int phase, const std::vector& local_token_nums); + + /// Zero-fills captured rows on warmup rounds, where no capture happens. + void SeedWarmup(int phase, cudaStream_t stream); + + /// Reunites the owned rows of `local_buffer` (a per-phase backing buffer of + /// capacity rows); returns the full-width active view the draft pass uses. + Tensor Gather(int phase, const Tensor& local_buffer, cudaStream_t stream); + + Tensor& captured(int phase) + { + return data_.at(phase).captured; + } + + int capacity() const + { + return capacity_; + } + + const comm::OwnedTokenRows& owned(int phase) const + { + return data_.at(phase).owned; + } + +private: + struct Data { + Tensor captured; + Tensor gathered_padded; + comm::OwnedTokenRows owned; + int token_num{}; + }; + + int capture_width_; + int hidden_units_; + int capacity_; + int attn_dp_rank_; + int model_tp_group_; + int model_tp_rank_; + int model_tp_size_; + comm::DeviceCommImpl* d_comm_; + std::vector data_; +}; + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/eagle3/eagle3_model.cc b/src/turbomind/models/speculative/eagle3/eagle3_model.cc new file mode 100644 index 0000000000..e1512a9f33 --- /dev/null +++ b/src/turbomind/models/speculative/eagle3/eagle3_model.cc @@ -0,0 +1,159 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#include "src/turbomind/models/speculative/eagle3/eagle3_model.h" + +#include "src/turbomind/comm/device_comm.h" +#include "src/turbomind/comm/token_ownership.h" +#include "src/turbomind/core/check.h" +#include "src/turbomind/core/context.h" +#include "src/turbomind/kernels/draft_carry_kernels.h" +#include "src/turbomind/kernels/norm/rms_norm.h" +#include "src/turbomind/models/decoder_layer_weight.h" +#include "src/turbomind/models/model_weight.h" +#include "src/turbomind/models/speculative/eagle3/eagle3_weight.h" +#include "src/turbomind/models/speculative/eagle3/target_hidden_projection.h" +#include "src/turbomind/models/speculative/registry.h" + +namespace turbomind { + +Eagle3Model::Eagle3Model(const SpeculativeModelArgs& args): + FixedChainSpeculativeModel(args), + spec_weights_{*TM_CHECK_NOTNULL(args.draft_weights.get("spec"))}, + data_(args.phases) +{ + const auto& p = args.param; + const auto& dw = draft_.weights(); + + for (Data& d : data_) { + d.draft_attention_input = Tensor{{p.max_forward_token_num, 2 * dw.hidden_units}, dw.data_type, kDEVICE}; + + if (comm_.d_comm) { + auto symmetric_allocator = GetSymmAllocator(comm_.d_comm); + d.draft_carry = Tensor{{p.max_batch_size, dw.hidden_units}, dw.data_type, symmetric_allocator}; + } + else { + d.draft_carry = Tensor{{p.max_batch_size, dw.hidden_units}, dw.data_type, kDEVICE}; + } + } + + projection_ = std::make_unique(p, + args.ctx, + args.phases, + target_.weights().num_layer, + target_.weights().hidden_units, + target_.weights().data_type, + p.spec_tap_layer_ids, + *spec_weights_.target_hidden_proj); +} + +Eagle3Model::~Eagle3Model() = default; + +HiddenStateTap* Eagle3Model::TapSource() +{ + return projection_.get(); +} + +Tensor Eagle3Model::InitialCarry(int phase, cudaStream_t stream) +{ + return projection_->ProjectAndGather(phase, stream); +} + +Tensor Eagle3Model::Embed(int, const Buffer_& ids, int rows, const DraftContext& ctx, TensorMap& env, EmbedStage) +{ + const auto& dw = draft_.weights(); + + Tensor storage = use_ag2d_ ? Tensor{env.at("symm_buf").buffer().view(dw.data_type), {rows, dw.hidden_units}} + : ctx.embedding_storage.slice({0, 0}, {rows, dw.hidden_units}); + return draft_.Embed(ids, storage, env); +} + +FixedChainSpeculativeModel::CombineResult Eagle3Model::Combine(int phase, Tensor embeddings, Tensor carry, int rows, TensorMap&) +{ + Data& data = data_[phase]; + const auto& dw = draft_.weights(); + const cudaStream_t stream = core::Context::stream().handle(); + + DecoderLayerWeight* const draft_layer = dw.layer(0); + const NormWeight* const draft_hidden_norm = TM_CHECK_NOTNULL(spec_weights_.hidden_norm(0)); + + Tensor attention_input = data.draft_attention_input.slice({0, 0}, {rows, 2 * dw.hidden_units}); + invokeRMSNormConcat(attention_input, + embeddings, + draft_layer->attention_norm->weight, + draft_layer->attention_norm->norm_eps_, + draft_layer->attention_norm->zero_centered_, + carry, + draft_hidden_norm->weight, + draft_hidden_norm->norm_eps_, + draft_hidden_norm->zero_centered_, + stream); + + CombineResult result; + result.residual = std::move(carry); + result.attention_input = std::move(attention_input); + return result; +} + +Tensor Eagle3Model::NextCarry(const LanguageModel::DecoderOutputs& out, + int step, + int phase, + TensorMap& env, + cudaStream_t stream) +{ + Data& data = data_[phase]; + FixedChainPhaseData& chain = fixed_chain_.phase_data(phase); + CommonData& common_data = common(phase); + const auto& dw = draft_.weights(); + + const int candidate_count = chain.draft_extension_query_count; + + Tensor carry = data.draft_carry.slice({0, 0}, {candidate_count, dw.hidden_units}); + if (candidate_count == 0) { + return carry; + } + + const auto& batch_local_token_nums = env.at("batch").data()[0]->local_token_num; + + const int global_rank = comm_.d_comm ? comm_.d_comm->rank(0) : 0; + const int tp0_size = comm_.d_comm ? comm_.d_comm->n_ranks(0) : 1; + const int tp1_size = comm_.d_comm ? comm_.d_comm->n_ranks(comm_.d_tp_group) : 1; + + const comm::OwnedTokenRows owned = step == 0 ? + comm::ComputeTokenOwnership(global_rank, tp0_size, tp1_size, batch_local_token_nums.data()) : + comm::ComputeTokenOwnership(global_rank, tp0_size, tp1_size, chain.draft_extension_local_token_nums.data()); + + const Buffer_ selected_rows = step == 0 ? common_data.draft_selected_token_pos.slice(0, candidate_count) : + draft_identity_token_pos_.slice(0, candidate_count); + + invokeSelectDraftCarry(out.pre_final_residual.raw_data(), + selected_rows.data(), + common_data.draft_candidate_active.data(), + carry.raw_data(), + static_cast(out.pre_final_residual.shape(0)), + candidate_count, + dw.hidden_units, + byte_size(dw.data_type) * 8, + owned.local_begin(), + owned.local_end(), + stream); + + if (tp1_size > 1) { + comm_.d_comm->AllReduceSum(carry.raw_data(), + carry.raw_data(), + carry.size(), + carry.dtype(), + comm_.d_tp_group, + stream); + } + + return carry; +} + +LanguageModel& Eagle3Model::HeadModel() +{ + return draft_; +} + +TM_REGISTER_SPECULATIVE_MODEL("eagle3", Eagle3Model); + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/eagle3/eagle3_model.h b/src/turbomind/models/speculative/eagle3/eagle3_model.h new file mode 100644 index 0000000000..3305516c5d --- /dev/null +++ b/src/turbomind/models/speculative/eagle3/eagle3_model.h @@ -0,0 +1,48 @@ +// Copyright (c) OpenMMLab. All rights reserved. +#pragma once + +#include "src/turbomind/models/speculative/fixed_chain_model.h" + +#include +#include + +namespace turbomind { + +class Eagle3Weight; +class TargetHiddenProjection; + +class Eagle3Model final: public FixedChainSpeculativeModel { +public: + explicit Eagle3Model(const SpeculativeModelArgs& args); + ~Eagle3Model() override; + +protected: + HiddenStateTap* TapSource() override; + Tensor InitialCarry(int phase, cudaStream_t stream) override; + Tensor Embed(int phase, + const Buffer_& ids, + int rows, + const DraftContext& ctx, + TensorMap& env, + EmbedStage stage) override; + CombineResult Combine(int phase, Tensor embeddings, Tensor carry, int rows, TensorMap& env) override; + Tensor NextCarry(const LanguageModel::DecoderOutputs& out, + int step, + int phase, + TensorMap& env, + cudaStream_t stream) override; + LanguageModel& HeadModel() override; + +private: + struct Data { + Tensor draft_attention_input; + Tensor draft_carry; + }; + + const Eagle3Weight& spec_weights_; + + std::unique_ptr projection_; + std::vector data_; +}; + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/eagle3/eagle3_weight.cc b/src/turbomind/models/speculative/eagle3/eagle3_weight.cc new file mode 100644 index 0000000000..25ae7972a4 --- /dev/null +++ b/src/turbomind/models/speculative/eagle3/eagle3_weight.cc @@ -0,0 +1,32 @@ +// Copyright (c) OpenMMLab. All rights reserved. +#include "src/turbomind/models/speculative/eagle3/eagle3_weight.h" + +#include "src/turbomind/core/check.h" +#include "src/turbomind/core/registry.h" + +namespace turbomind { + +bool Eagle3Weight::verify(std::vector& missing) +{ + Module::verify(missing); + if (!target_hidden_proj) { + missing.push_back(full_path() + ": missing target_hidden_proj"); + } + if (!hidden_norms || hidden_norms->size() == 0) { + missing.push_back(full_path() + ": missing hidden_norms"); + } + return missing.empty(); +} + +NormWeight* Eagle3Weight::hidden_norm(int layer) const +{ + if (!hidden_norms) { + return nullptr; + } + return static_cast(hidden_norms->child(std::to_string(layer))); +} + +TM_MODULE_REGISTER(Eagle3Weight, core::Eagle3WeightConfig); +TM_MODULE_METHODS(Eagle3Weight, EAGLE3_WEIGHT_CHILDREN, EAGLE3_WEIGHT_PARAMS) + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/eagle3/eagle3_weight.h b/src/turbomind/models/speculative/eagle3/eagle3_weight.h new file mode 100644 index 0000000000..3dad60465c --- /dev/null +++ b/src/turbomind/models/speculative/eagle3/eagle3_weight.h @@ -0,0 +1,54 @@ +// Copyright (c) OpenMMLab. All rights reserved. +#pragma once + +#include "src/turbomind/core/module.h" +#include "src/turbomind/models/linear_weight.h" +#include "src/turbomind/models/norm_weight.h" + +#include +#include + +namespace turbomind::core { + +struct Eagle3WeightConfig: ModuleConfig { + Eagle3WeightConfig(): ModuleConfig{"Eagle3Weight"} {} + +#define EAGLE3_WEIGHT_FIELDS(X) X(DataType, data_type) + + EAGLE3_WEIGHT_FIELDS(TM_MEMBER) + TM_FOR_EACH(Eagle3WeightConfig, EAGLE3_WEIGHT_FIELDS) + +#undef EAGLE3_WEIGHT_FIELDS +}; + +} // namespace turbomind::core + +namespace turbomind { + +/// EAGLE3's own weight tree. Attached at ModelWeight::spec on the draft tree. +class Eagle3Weight: public core::Module { +public: + const char* type() const override + { + return "Eagle3Weight"; + } + + Eagle3Weight() = default; + + explicit Eagle3Weight(const core::Eagle3WeightConfig&) {} + + bool verify(std::vector& missing) override; + + /// Per-draft-layer hidden norm; null when the layer has none. + NormWeight* hidden_norm(int layer) const; + +#define EAGLE3_WEIGHT_CHILDREN(X) \ + X(LinearWeight, target_hidden_proj) \ + X(core::ModuleList, hidden_norms) + +#define EAGLE3_WEIGHT_PARAMS(X) + + TM_MODULE_DECLARE(Eagle3Weight, EAGLE3_WEIGHT_CHILDREN, EAGLE3_WEIGHT_PARAMS) +}; + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/eagle3/target_hidden_projection.cc b/src/turbomind/models/speculative/eagle3/target_hidden_projection.cc new file mode 100644 index 0000000000..e6d9efd773 --- /dev/null +++ b/src/turbomind/models/speculative/eagle3/target_hidden_projection.cc @@ -0,0 +1,117 @@ +#include "src/turbomind/models/speculative/eagle3/target_hidden_projection.h" + +#include +#include + +#include "src/turbomind/core/context.h" +#include "src/turbomind/core/scope.h" +#include "src/turbomind/kernels/core/math.h" +#include "src/turbomind/models/linear_weight.h" +#include "src/turbomind/models/llama/LlamaLinear.h" +#include "src/turbomind/models/llama/context.h" + +namespace turbomind { + +TargetHiddenProjection::TargetHiddenProjection(const EngineParam& engine, + const Context& context, + int phases, + int target_layer_count, + int hidden_units, + DataType data_type, + std::vector target_layer_ids, + const LinearWeight& projection_weight): + HiddenStateTap{*context.is_warm_up}, + data_type_{data_type}, + hidden_units_{hidden_units}, + target_layer_count_{target_layer_count}, + linear_{*context.linear}, + projection_weight_{projection_weight}, + target_layer_ids_{std::move(target_layer_ids)}, + tap_ordinal_by_completed_layer_(target_layer_count + 1, -1), + collect_{engine, + context, + phases, + static_cast(target_layer_ids_.size()) * hidden_units, + hidden_units, + data_type}, + data_(phases) +{ + for (int ordinal = 0; ordinal < static_cast(target_layer_ids_.size()); ++ordinal) { + tap_ordinal_by_completed_layer_[target_layer_ids_[ordinal]] = ordinal; + } + + for (auto& d : data_) { + d.projected_local = Tensor{{collect_.capacity(), hidden_units_}, data_type_, kDEVICE}; + } +} + +void TargetHiddenProjection::Begin(int phase, const std::vector& local_token_nums) +{ + BeginPhase(phase); + collect_.Begin(phase, local_token_nums); +} + +void TargetHiddenProjection::BeginPhase(int phase) +{ + active_phase_ = phase; +} + +void TargetHiddenProjection::SeedWarmup(int phase, cudaStream_t stream) +{ + collect_.SeedWarmup(phase, stream); +} + +int TargetHiddenProjection::TapOrdinal(int completed_layer_count) const +{ + return tap_ordinal_by_completed_layer_[completed_layer_count]; +} + +void TargetHiddenProjection::Capture(int phase, int tap_ordinal, const Tensor& local_residual, cudaStream_t stream) +{ + const comm::OwnedTokenRows& owned = collect_.owned(phase); + Tensor& captured = collect_.captured(phase); + const int first = owned.local_begin(); + const int rows = owned.row_count(); + + if (rows == 0) { + return; + } + + invokeCaptureTargetHiddenRows(local_residual.raw_data(), + captured.raw_data(), + first, + rows, + hidden_units_, + local_residual.stride(0), + captured.stride(0), + tap_ordinal, + byte_size(data_type_) * 8, + stream); +} + +Tensor TargetHiddenProjection::ProjectLocal(int phase) +{ + Data& d = data_[phase]; + const int tap_count = static_cast(target_layer_ids_.size()); + const int input_width = tap_count * hidden_units_; + const int rows = collect_.owned(phase).row_count(); + + Tensor input = collect_.captured(phase).slice({0, 0}, {rows, input_width}); + Tensor output = d.projected_local.slice({0, 0}, {rows, hidden_units_}); + + if (rows == 0) { + return output; + } + + TM_SCOPE_CALL(linear_.Forward(input, projection_weight_, output)); + + return output; +} + +Tensor TargetHiddenProjection::ProjectAndGather(int phase, cudaStream_t stream) +{ + Tensor active = ProjectLocal(phase); + return collect_.Gather(phase, data_[phase].projected_local, stream); +} + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/eagle3/target_hidden_projection.h b/src/turbomind/models/speculative/eagle3/target_hidden_projection.h new file mode 100644 index 0000000000..bd447e4258 --- /dev/null +++ b/src/turbomind/models/speculative/eagle3/target_hidden_projection.h @@ -0,0 +1,70 @@ +#pragma once + +#include + +#include "src/turbomind/core/core.h" +#include "src/turbomind/models/speculative/collect_hidden_states.h" +#include "src/turbomind/models/speculative/eagle3/target_hidden_projection_kernels.h" +#include "src/turbomind/models/speculative/hidden_state_tap.h" + +namespace turbomind { + +class Context; +class EngineParam; +class LinearWeight; +class LlamaLinear; + +/// EAGLE3's tap: captures residuals of the configured target layers, projects +/// the concatenated rows through the eagle fc, and reunites the result across +/// model-TP ranks for the draft pass. +class TargetHiddenProjection: public HiddenStateTap { +public: + TargetHiddenProjection(const EngineParam& engine, + const Context& context, + int phases, + int target_layer_count, + int hidden_units, + DataType data_type, + std::vector target_layer_ids, + const LinearWeight& projection_weight); + + void Begin(int phase, const std::vector& local_token_nums) override; + + void BeginPhase(int phase); + + void SeedWarmup(int phase, cudaStream_t stream) override; + + int TapOrdinal(int completed_layer_count) const override; + + void Capture(int tap_ordinal, + const Tensor& local_residual, + const Tensor&, + cudaStream_t stream) override + { + Capture(active_phase_, tap_ordinal, local_residual, stream); + } + + void Capture(int phase, int tap_ordinal, const Tensor& local_residual, cudaStream_t stream); + + Tensor ProjectLocal(int phase); + + Tensor ProjectAndGather(int phase, cudaStream_t stream); + +private: + struct Data { + Tensor projected_local; + }; + + DataType data_type_; + int hidden_units_; + int target_layer_count_; + LlamaLinear& linear_; + const LinearWeight& projection_weight_; + std::vector target_layer_ids_; + std::vector tap_ordinal_by_completed_layer_; + CollectHiddenStates collect_; + std::vector data_; + int active_phase_{}; +}; + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/eagle3/target_hidden_projection_kernels.cu b/src/turbomind/models/speculative/eagle3/target_hidden_projection_kernels.cu new file mode 100644 index 0000000000..565d9e554b --- /dev/null +++ b/src/turbomind/models/speculative/eagle3/target_hidden_projection_kernels.cu @@ -0,0 +1,35 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#include "src/turbomind/models/speculative/eagle3/target_hidden_projection_kernels.h" + +#include + +namespace turbomind { + +void invokeCaptureTargetHiddenRows(const void* packed_residual, + void* captured, + int owned_begin, + int owned_row_count, + int hidden_units, + int packed_leading_dimension, + int captured_leading_dimension, + int tap_ordinal, + int element_bits, + cudaStream_t stream) +{ + if (owned_row_count == 0) { + return; + } + + cudaMemcpy2DAsync(static_cast(captured) + tap_ordinal * hidden_units * element_bits / 8, + captured_leading_dimension * element_bits / 8, + static_cast(packed_residual) + + owned_begin * packed_leading_dimension * element_bits / 8, + packed_leading_dimension * element_bits / 8, + hidden_units * element_bits / 8, + owned_row_count, + cudaMemcpyDeviceToDevice, + stream); +} + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/eagle3/target_hidden_projection_kernels.h b/src/turbomind/models/speculative/eagle3/target_hidden_projection_kernels.h new file mode 100644 index 0000000000..806bc979e0 --- /dev/null +++ b/src/turbomind/models/speculative/eagle3/target_hidden_projection_kernels.h @@ -0,0 +1,20 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#pragma once + +#include + +namespace turbomind { + +void invokeCaptureTargetHiddenRows(const void* packed_residual, + void* captured, + int owned_begin, + int owned_row_count, + int hidden_units, + int packed_leading_dimension, + int captured_leading_dimension, + int tap_ordinal, + int element_bits, + cudaStream_t stream); + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/eagle3/target_hidden_projection_python_bind.cpp b/src/turbomind/models/speculative/eagle3/target_hidden_projection_python_bind.cpp new file mode 100644 index 0000000000..1fb77b0d46 --- /dev/null +++ b/src/turbomind/models/speculative/eagle3/target_hidden_projection_python_bind.cpp @@ -0,0 +1,59 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#include + +#include + +#include + +#include "src/turbomind/models/speculative/eagle3/target_hidden_projection_kernels.h" +#include "src/turbomind/python/eagle3_component_bindings.h" +#include "src/turbomind/python/eagle3_dlpack_internal.h" +#include "src/turbomind/utils/cuda_utils.h" + +namespace py = pybind11; + +namespace turbomind::python { +namespace { + +int GetCudaOrdinal(py::handle tensor) +{ + return tensor.attr("__dlpack_device__")().cast()[1].cast(); +} + +} // namespace + +void BindTargetHiddenProjection(py::module_& module) +{ + module.def( + "capture_target_hidden_rows", + [](py::handle packed_residual_object, + py::handle captured_object, + int owned_begin, + int owned_row_count, + int tap_ordinal, + uintptr_t stream_ptr) { + CudaDeviceGuard guard{GetCudaOrdinal(packed_residual_object)}; + auto packed_residual = detail::ConsumeDLPackWithStrides(packed_residual_object, stream_ptr); + auto captured = detail::ConsumeDLPackWithStrides(captured_object, stream_ptr); + + invokeCaptureTargetHiddenRows(packed_residual.data_or(static_cast(nullptr)), + captured.data_or(static_cast(nullptr)), + owned_begin, + owned_row_count, + static_cast(packed_residual.shape(1)), + static_cast(packed_residual.stride(0)), + static_cast(captured.stride(0)), + tap_ordinal, + static_cast(byte_size(packed_residual.dtype(), 8)), + reinterpret_cast(stream_ptr)); + }, + py::arg("packed_residual"), + py::arg("captured"), + py::arg("owned_begin"), + py::arg("owned_row_count"), + py::arg("tap_ordinal"), + py::arg("stream_ptr")); +} + +} // namespace turbomind::python diff --git a/src/turbomind/models/speculative/fixed_chain_model.cc b/src/turbomind/models/speculative/fixed_chain_model.cc new file mode 100644 index 0000000000..2d60e97d7c --- /dev/null +++ b/src/turbomind/models/speculative/fixed_chain_model.cc @@ -0,0 +1,233 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#include "src/turbomind/models/speculative/fixed_chain_model.h" + +#include + +#include "src/turbomind/comm/device_comm.h" +#include "src/turbomind/core/context.h" +#include "src/turbomind/core/copy.h" +#include "src/turbomind/kernels/speculative_sequence_kernels.h" +#include "src/turbomind/models/model_weight.h" +#include "src/turbomind/models/speculative/hidden_state_tap.h" + +namespace turbomind { + +FixedChainSpeculativeModel::FixedChainSpeculativeModel(const SpeculativeModelArgs& args): + target_{args.target}, + comm_{args.ctx.comm}, + use_ag2d_{comm_.d_comm && comm_.d_comm->Query(comm::kHasAllGather2D)}, + draft_{args.registry, args.param, args.ctx, args.draft_weights, args.phases}, + draft_hidden_{draft_.weights().hidden_units}, + policy_{args.param.spec_num_draft_tokens}, + fixed_chain_{args.ctx.comm, args.phases, args.param.max_batch_size}, + data_(args.phases) +{ + const EngineParam& param = args.param; + + Buffer_ identity_host{param.max_batch_size, kCPUpinned}; + draft_identity_token_pos_ = {param.max_batch_size, kDEVICE}; + std::iota(identity_host.data(), identity_host.data() + param.max_batch_size, 0); + Copy(identity_host, draft_identity_token_pos_); + core::Context::stream().Sync(); + + const DataType draft_type = draft_.weights().data_type; + + for (CommonData& data : data_) { + data.draft_input_ids = {param.max_forward_token_num, kDEVICE}; + data.draft_selected_token_pos = {param.max_batch_size, kDEVICE}; + data.draft_candidate_active = {param.max_batch_size, kDEVICE}; + data.draft_proposal_ids = {param.max_batch_size, kDEVICE}; + + data.draft_selected_normalized_hidden = { + {param.max_batch_size, draft_hidden_}, draft_type, kDEVICE}; + } +} + +FixedChainSpeculativeModel::~FixedChainSpeculativeModel() = default; + +const SpeculativePolicy& FixedChainSpeculativeModel::policy() const +{ + return policy_; +} + +HiddenStateTap* FixedChainSpeculativeModel::Tap(int) +{ + return TapSource(); +} + +void FixedChainSpeculativeModel::Setup(int phase, TensorMap& env) +{ + fixed_chain_.Setup(phase, + env.at("requests").buffer(), + *env.at("copy").data()[0]); + draft_.Run(BatchOp::kSetup, phase, env); +} + +void FixedChainSpeculativeModel::Run(BatchOp op, int phase, TensorMap& env) +{ + if (op == BatchOp::kAdd) { + draft_.Run(op, phase, env); + return; + } + if (op == BatchOp::kSetup) { + Setup(phase, env); + return; + } + if (op == BatchOp::kPrepare) { + TensorMap draft_env = env; + draft_env.at("finished") = env.at("finished_on_entry"); + draft_env.at("q_offsets") = env.at("q_offsets"); + draft_env.at("k_offsets") = env.at("k_offsets"); + draft_.Run(BatchOp::kPrepare, phase, draft_env); + } +} + +void FixedChainSpeculativeModel::RunDraft(int phase, const DraftContext& ctx, TensorMap& env) +{ + CommonData& data = data_[phase]; + FixedChainPhaseData& chain = fixed_chain_.phase_data(phase); + + const int target_query_count = chain.refresh_decode.query_count + chain.refresh_prefill.query_count; + const int candidate_count = chain.draft_extension_query_count; + const int k = policy_.max_proposals(); + const cudaStream_t stream = core::Context::stream().handle(); + + AttentionForwardMetadata refresh_metadata{}; + refresh_metadata.decode = chain.refresh_decode; + refresh_metadata.prefill = chain.refresh_prefill; + refresh_metadata.q_offsets = ctx.target_q_offsets; + refresh_metadata.k_offsets = ctx.target_k_offsets; + + AttentionForwardMetadata extension_metadata{}; + extension_metadata.decode = chain.extension_decode; + extension_metadata.prefill = {}; + extension_metadata.q_offsets = chain.draft_extension_q_offsets.slice(0, ctx.batch_size + 1); + extension_metadata.k_offsets = chain.draft_extension_k_offsets.slice(0, ctx.batch_size + 1); + + Tensor extension_local_token_nums{chain.draft_extension_local_token_nums.data(), + Layout{{static_cast(chain.draft_extension_local_token_nums.size())}}, + kCPU}; + + const bool run_extensions = k > 1 && chain.draft_extension_global_query_count > 0; + + Tensor carry = InitialCarry(phase, stream); + + invokeBuildDraftRefreshInputs(data.draft_input_ids.data(), + data.draft_selected_token_pos.data(), + data.draft_candidate_active.data(), + ctx.request_token_ids_ptrs, + ctx.target_q_offsets.data(), + ctx.target_k_offsets.data(), + chain.draft_extension_q_offsets.data(), + ctx.accept_len.data(), + chain.limit_to_accept_len.data(), + env.at("finished").data(), + target_query_count, + ctx.batch_size, + candidate_count, + stream); + + Tensor embeddings = Embed(phase, + data.draft_input_ids.slice(0, target_query_count), + target_query_count, + ctx, + env, + EmbedStage::kRefresh); + + TensorMap draft_env = env; + draft_env.at("finished") = ctx.finished_on_entry; + draft_env.at("q_offsets") = ctx.target_q_offsets; + draft_env.at("k_offsets") = ctx.target_k_offsets; + + CombineResult combined = Combine(phase, std::move(embeddings), std::move(carry), target_query_count, env); + + LanguageModel::DecoderInputs refresh_in; + refresh_in.residual = std::move(combined.residual); + refresh_in.attention_input = std::move(combined.attention_input); + refresh_in.selected_token_pos = data.draft_selected_token_pos.slice(0, candidate_count); + refresh_in.selected_hidden_buffer = data.draft_selected_normalized_hidden.slice( + {0, 0}, {candidate_count, draft_hidden_}); + refresh_in.attention_metadata = &refresh_metadata; + + LanguageModel::DecoderOutputs out = draft_.RunDecoder(phase, refresh_in, draft_env); + + if (candidate_count > 0) { + const auto& hw = HeadModel().weights(); + Tensor logits = ctx.head_storage.slice({0, 0}, {candidate_count, hw.vocab_size_padded}); + logits = HeadModel().Logits(out.selected_hidden, logits, draft_env); + invokeDraftArgmaxAndStoreToken(logits, + data.draft_proposal_ids.data(), + ctx.request_token_ids_ptrs, + chain.draft_extension_q_offsets.data(), + data.draft_candidate_active.data(), + ctx.sequence_length.data(), + ctx.accept_len.data(), + ctx.batch_size, + candidate_count, + 0, + hw.vocab_size, + stream); + } + + if (!run_extensions) { + return; + } + + TensorMap extension_env = env; + extension_env.at("finished") = ctx.finished_on_entry; + extension_env.at("q_offsets") = chain.draft_extension_q_offsets.slice(0, ctx.batch_size + 1); + extension_env.at("k_offsets") = chain.draft_extension_k_offsets.slice(0, ctx.batch_size + 1); + extension_env.produce("decoder_local_token_nums", extension_local_token_nums); + + for (int step = 1; step < k; ++step) { + const int i = step - 1; + + carry = NextCarry(out, i, phase, env, stream); + embeddings = Embed(phase, + data.draft_proposal_ids.slice(0, candidate_count), + candidate_count, + ctx, + extension_env, + EmbedStage::kExtension); + combined = Combine(phase, std::move(embeddings), std::move(carry), candidate_count, extension_env); + + invokeBuildDraftExtensionKeyOffsets(chain.draft_extension_k_offsets.data(), + chain.draft_extension_q_offsets.data(), + ctx.sequence_length.data(), + ctx.accept_len.data(), + ctx.batch_size, + i, + stream); + + LanguageModel::DecoderInputs extension_in; + extension_in.residual = std::move(combined.residual); + extension_in.attention_input = std::move(combined.attention_input); + extension_in.selected_token_pos = draft_identity_token_pos_.slice(0, candidate_count); + extension_in.selected_hidden_buffer = data.draft_selected_normalized_hidden.slice( + {0, 0}, {candidate_count, draft_hidden_}); + extension_in.attention_metadata = &extension_metadata; + + out = draft_.RunDecoder(phase, extension_in, extension_env); + + if (candidate_count > 0) { + const auto& hw = HeadModel().weights(); + Tensor logits = ctx.head_storage.slice({0, 0}, {candidate_count, hw.vocab_size_padded}); + logits = HeadModel().Logits(out.selected_hidden, logits, extension_env); + invokeDraftArgmaxAndStoreToken(logits, + data.draft_proposal_ids.data(), + ctx.request_token_ids_ptrs, + chain.draft_extension_q_offsets.data(), + data.draft_candidate_active.data(), + ctx.sequence_length.data(), + ctx.accept_len.data(), + ctx.batch_size, + candidate_count, + i + 1, + hw.vocab_size, + stream); + } + } +} + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/fixed_chain_model.h b/src/turbomind/models/speculative/fixed_chain_model.h new file mode 100644 index 0000000000..1c8a59ceac --- /dev/null +++ b/src/turbomind/models/speculative/fixed_chain_model.h @@ -0,0 +1,100 @@ +// Copyright (c) OpenMMLab. All rights reserved. +#pragma once + +#include "src/turbomind/models/language_model.h" +#include "src/turbomind/models/llama/context.h" +#include "src/turbomind/models/llama/unified_attention_layer.h" +#include "src/turbomind/models/speculative/fixed_chain_policy.h" +#include "src/turbomind/models/speculative/fixed_chain_setup.h" +#include "src/turbomind/models/speculative/speculative_model.h" + +#include + +namespace turbomind { + +/// Shared skeleton for fixed-chain speculators. Implements the draft pass as one +/// refresh step followed by k - 1 extension steps, and delegates the per-model +/// math (embedding table, carry, norm + projection, LM head) to the hooks below. +/// See docs/adr/0001-fixed-chain-speculative-model-base.md. +class FixedChainSpeculativeModel: public SpeculativeModel { +public: + explicit FixedChainSpeculativeModel(const SpeculativeModelArgs& args); + ~FixedChainSpeculativeModel() override; + + const SpeculativePolicy& policy() const override; + + void Run(BatchOp op, int phase, TensorMap& env) override; + + HiddenStateTap* Tap(int phase) override; + + void RunDraft(int phase, const DraftContext& ctx, TensorMap& env) final; + +protected: + struct CommonData { + Buffer_ draft_input_ids; + Buffer_ draft_selected_token_pos; + Buffer_ draft_candidate_active; + Buffer_ draft_proposal_ids; + Tensor draft_selected_normalized_hidden; + }; + + enum class EmbedStage { kRefresh, kExtension }; + + struct CombineResult { + Tensor residual; + Tensor attention_input; + }; + + /// The tap whose captured state feeds the draft pass. + virtual HiddenStateTap* TapSource() = 0; + + /// Carry entering the refresh step: the target's tapped state, collected. + virtual Tensor InitialCarry(int phase, cudaStream_t stream) = 0; + + /// Embedding lookup for the token entering the current step. `rows` sizes the + /// storage; refresh-stage hooks may patch the embeddings in place. + virtual Tensor Embed(int phase, + const Buffer_& ids, + int rows, + const DraftContext& ctx, + TensorMap& env, + EmbedStage stage) = 0; + + /// Norm-concat of embeddings and carry, projected to the decoder input. The + /// returned residual feeds the decoder; attention_input is optional. + virtual CombineResult Combine(int phase, Tensor embeddings, Tensor carry, int rows, TensorMap& env) = 0; + + /// Carry entering step `step + 1`, derived from the previous decoder output. + /// `step` is 0 when `out` is the refresh output. + virtual Tensor NextCarry(const LanguageModel::DecoderOutputs& out, + int step, + int phase, + TensorMap& env, + cudaStream_t stream) = 0; + + /// The language model whose LM head turns draft hiddens into logits. + virtual LanguageModel& HeadModel() = 0; + + CommonData& common(int phase) + { + return data_.at(phase); + } + + LanguageModel& target_; + LanguageModel draft_; + const Communicators& comm_; + const bool use_ag2d_; + const int draft_hidden_; + + FixedChainPolicy policy_; + FixedChainSetup fixed_chain_; + + Buffer_ draft_identity_token_pos_; + +private: + void Setup(int phase, TensorMap& env); + + std::vector data_; +}; + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/fixed_chain_policy.cc b/src/turbomind/models/speculative/fixed_chain_policy.cc new file mode 100644 index 0000000000..584428591c --- /dev/null +++ b/src/turbomind/models/speculative/fixed_chain_policy.cc @@ -0,0 +1,27 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#include "src/turbomind/models/speculative/fixed_chain_policy.h" + +namespace turbomind { + +int FixedChainPolicy::max_proposals() const +{ + return proposal_count_; +} + +int FixedChainPolicy::token_row_tail() const +{ + return 2; +} + +RoundExtent FixedChainPolicy::Extent(const RoundRequest&) const +{ + return {proposal_count_ + 1, proposal_count_ - 1}; +} + +std::optional FixedChainPolicy::Bootstrap(int prompt_len) const +{ + return BootstrapExtent{prompt_len + proposal_count_ - 1, prompt_len + 2 * proposal_count_}; +} + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/fixed_chain_policy.h b/src/turbomind/models/speculative/fixed_chain_policy.h new file mode 100644 index 0000000000..a1813e9d73 --- /dev/null +++ b/src/turbomind/models/speculative/fixed_chain_policy.h @@ -0,0 +1,26 @@ +// Copyright (c) OpenMMLab. All rights reserved. +#pragma once + +#include "src/turbomind/models/speculative/speculative_policy.h" + +namespace turbomind { + +/// Fixed-width chain drafting: k proposals are verified as k + 1 query rows, +/// with a k - 1 private cache tail for the next proposal round. +class FixedChainPolicy final: public SpeculativePolicy { +public: + explicit FixedChainPolicy(int proposal_count): proposal_count_{proposal_count} {} + + int max_proposals() const override; + + int token_row_tail() const override; + + RoundExtent Extent(const RoundRequest& request) const override; + + std::optional Bootstrap(int prompt_len) const override; + +private: + const int proposal_count_; +}; + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/fixed_chain_setup.cc b/src/turbomind/models/speculative/fixed_chain_setup.cc new file mode 100644 index 0000000000..cb28233022 --- /dev/null +++ b/src/turbomind/models/speculative/fixed_chain_setup.cc @@ -0,0 +1,90 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#include "src/turbomind/models/speculative/fixed_chain_setup.h" + +#include "src/turbomind/comm/host_comm.h" +#include "src/turbomind/core/copy.h" +#include "src/turbomind/engine/request.h" + +#include +#include + +namespace turbomind { + +FixedChainSetup::FixedChainSetup(const Communicators& comm, int phases, int max_batch_size): + comm_{comm}, limit_to_accept_len_host_{max_batch_size, kCPUpinned}, data_(phases) +{ + for (FixedChainPhaseData& data : data_) { + data.draft_extension_q_offsets_host = {max_batch_size + 1, kCPUpinned}; + data.draft_extension_q_offsets = {max_batch_size + 1, kDEVICE}; + data.draft_extension_k_offsets = {max_batch_size + 1, kDEVICE}; + data.draft_extension_local_token_nums.assign(comm_.h_dp_group->n_ranks(), 0); + data.limit_to_accept_len = {max_batch_size, kDEVICE}; + } +} + +void FixedChainSetup::Setup(int phase, const Buffer_& requests, core::BatchCopy& copy) +{ + FixedChainPhaseData& data = data_.at(phase); + const int batch_size = requests.size(); + + Buffer_& extension = data.draft_extension_q_offsets_host; + extension[0] = 0; + + data.refresh_decode = {}; + data.refresh_prefill = {}; + + const int decode_request_count = + std::find_if(requests.begin(), requests.end(), [](const Sequence* request) { + return request->submitted->input_len > 1; + }) + - requests.begin(); + + for (int b = 0; b < batch_size; ++b) { + const SubmittedRow& row = *requests[b]->submitted; + + extension[b + 1] = extension[b] + static_cast(row.is_extension_candidate()); + limit_to_accept_len_host_[b] = row.autoregres; + + auto& partition = b < decode_request_count ? data.refresh_decode : data.refresh_prefill; + partition.request_count += 1; + partition.query_count += row.input_len; + partition.max_query_length = std::max(partition.max_query_length, row.input_len); + partition.key_capacity_sum += row.key_capacity_end; + partition.max_key_capacity = std::max(partition.max_key_capacity, row.key_capacity_end); + } + + data.draft_extension_query_count = extension[batch_size]; + + const int attn_dp_rank = comm_.h_dp_group->rank(); + std::fill(data.draft_extension_local_token_nums.begin(), data.draft_extension_local_token_nums.end(), 0); + data.draft_extension_local_token_nums[attn_dp_rank] = data.draft_extension_query_count; + + if (comm_.h_dp_group->n_ranks() > 1) { + comm::AllGather(comm_.h_dp_group, data.draft_extension_local_token_nums.data(), 1); + } + + data.draft_extension_global_query_count = + std::accumulate(data.draft_extension_local_token_nums.begin(), + data.draft_extension_local_token_nums.end(), + 0); + + data.extension_decode = {}; + data.extension_decode.request_count = batch_size; + data.extension_decode.query_count = data.draft_extension_query_count; + data.extension_decode.max_query_length = data.draft_extension_query_count ? 1 : 0; + + for (int b = 0; b < batch_size; ++b) { + if (extension[b + 1] == extension[b]) { + continue; + } + const int capacity = requests[b]->submitted->cache_write_end; + data.extension_decode.key_capacity_sum += capacity; + data.extension_decode.max_key_capacity = std::max(data.extension_decode.max_key_capacity, capacity); + } + + copy(extension, batch_size + 1, data.draft_extension_q_offsets); + copy(limit_to_accept_len_host_, batch_size, data.limit_to_accept_len); +} + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/fixed_chain_setup.h b/src/turbomind/models/speculative/fixed_chain_setup.h new file mode 100644 index 0000000000..712980e068 --- /dev/null +++ b/src/turbomind/models/speculative/fixed_chain_setup.h @@ -0,0 +1,51 @@ +// Copyright (c) OpenMMLab. All rights reserved. +#pragma once + +#include "src/turbomind/core/buffer.h" +#include "src/turbomind/models/llama/context.h" +#include "src/turbomind/models/llama/unified_attention_layer.h" + +#include + +namespace turbomind { + +class Sequence; + +namespace core { +class BatchCopy; +} + +struct FixedChainPhaseData { + Buffer_ draft_extension_q_offsets_host; + Buffer_ draft_extension_q_offsets; + Buffer_ draft_extension_k_offsets; + int draft_extension_query_count{}; + + std::vector draft_extension_local_token_nums; + int draft_extension_global_query_count{}; + + AttentionForwardMetadata::Partition refresh_decode; + AttentionForwardMetadata::Partition refresh_prefill; + AttentionForwardMetadata::Partition extension_decode; + + Buffer_ limit_to_accept_len; +}; + +class FixedChainSetup { +public: + FixedChainSetup(const Communicators& comm, int phases, int max_batch_size); + + void Setup(int phase, const Buffer_& requests, core::BatchCopy& copy); + + FixedChainPhaseData& phase_data(int phase) + { + return data_.at(phase); + } + +private: + const Communicators& comm_; + Buffer_ limit_to_accept_len_host_; + std::vector data_; +}; + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/hidden_state_tap.h b/src/turbomind/models/speculative/hidden_state_tap.h new file mode 100644 index 0000000000..608e93eeb2 --- /dev/null +++ b/src/turbomind/models/speculative/hidden_state_tap.h @@ -0,0 +1,54 @@ +// Copyright (c) OpenMMLab. All rights reserved. +#pragma once + +#include "src/turbomind/core/core.h" + +#include + +#include + +namespace turbomind { + +class HiddenStateTap { +public: + explicit HiddenStateTap(const int& is_warm_up): + is_warm_up_{is_warm_up} + { + } + + virtual ~HiddenStateTap() = default; + + /// Arms the tap for a target pass: records the token rows this rank owns + /// for the round and sizes the captured state. On warm-up rounds, where no + /// capture happens, zero-fills the captured state instead and returns null + /// so the decoder runs untapped. Arming is the tap's own business; callers + /// hold no warm-up knowledge. + HiddenStateTap* Arm(int phase, const std::vector& local_token_nums, cudaStream_t stream) + { + Begin(phase, local_token_nums); + if (is_warm_up_) { + SeedWarmup(phase, stream); + return nullptr; + } + return this; + } + + /// Called before each target forward: records the token rows this rank owns + /// for the round and sizes the captured state. + virtual void Begin(int phase, const std::vector& local_token_nums) = 0; + + /// Zero-fills captured state on warmup rounds, where no capture happens. + virtual void SeedWarmup(int phase, cudaStream_t stream) = 0; + + virtual int TapOrdinal(int completed_layer_count) const = 0; + + virtual void Capture(int tap_ordinal, + const Tensor& local_residual, + const Tensor& local_normalized_hidden, + cudaStream_t stream) = 0; + +private: + const int& is_warm_up_; +}; + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/qwen3_5_mtp/qwen3_5_mtp_model.cc b/src/turbomind/models/speculative/qwen3_5_mtp/qwen3_5_mtp_model.cc new file mode 100644 index 0000000000..0905a9589c --- /dev/null +++ b/src/turbomind/models/speculative/qwen3_5_mtp/qwen3_5_mtp_model.cc @@ -0,0 +1,179 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#include "src/turbomind/models/speculative/qwen3_5_mtp/qwen3_5_mtp_model.h" + +#include "src/turbomind/comm/device_comm.h" +#include "src/turbomind/core/context.h" +#include "src/turbomind/core/copy.h" +#include "src/turbomind/kernels/gpt_kernels.h" +#include "src/turbomind/kernels/norm/rms_norm.h" +#include "src/turbomind/models/input_processor.h" +#include "src/turbomind/models/model_weight.h" +#include "src/turbomind/models/speculative/qwen3_5_mtp/qwen3_5_mtp_weight.h" +#include "src/turbomind/models/speculative/qwen3_5_mtp/target_final_hidden.h" +#include "src/turbomind/models/speculative/registry.h" + +namespace turbomind { + +Qwen35MtpModel::Qwen35MtpModel(const SpeculativeModelArgs& args): + FixedChainSpeculativeModel(args), + spec_weights_{*args.draft_weights.get("spec")}, + linear_{*args.ctx.linear}, + hidden_units_{target_.weights().hidden_units}, + model_tp_rank_{args.param.model_tp_rank}, + model_tp_size_{args.param.attn_cp_size * args.param.attn_tp_size}, + final_hidden_{std::make_unique(args.param, + args.ctx, + args.phases, + target_.weights().num_layer, + hidden_units_, + target_.weights().data_type)}, + data_(args.phases) +{ + const EngineParam& param = args.param; + + Allocator projected_allocator = core::Context::device_alloc(); + if (model_tp_size_ > 1) { + projected_allocator = GetSymmAllocator(comm_.d_comm); + } + + for (Data& data : data_) { + data.normalized_concat = { + {param.max_forward_token_num, 2 * hidden_units_}, target_.weights().data_type, kDEVICE}; + data.projected_full = { + {param.max_forward_token_num, hidden_units_}, target_.weights().data_type, projected_allocator}; + } +} + +Qwen35MtpModel::~Qwen35MtpModel() = default; + +HiddenStateTap* Qwen35MtpModel::TapSource() +{ + return final_hidden_.get(); +} + +Tensor Qwen35MtpModel::InitialCarry(int phase, cudaStream_t stream) +{ + return final_hidden_->Gather(phase, stream); +} + +Tensor Qwen35MtpModel::Embed(int phase, + const Buffer_& ids, + int rows, + const DraftContext& ctx, + TensorMap& env, + EmbedStage stage) +{ + Tensor storage = use_ag2d_ ? Tensor{env.at("symm_buf").buffer().view(target_.weights().data_type), + {rows, hidden_units_}} + : ctx.embedding_storage.slice({0, 0}, {rows, hidden_units_}); + Tensor embeddings = target_.Embed(ids, storage, env); + + if (stage == EmbedStage::kRefresh) { + auto& copy = *env.at("copy").data()[0]; + ctx.input_processor->PatchSuccessorEmbedding(phase, embeddings, copy, env); + copy.Run(); + } + + return embeddings; +} + +FixedChainSpeculativeModel::CombineResult Qwen35MtpModel::Combine(int phase, + Tensor embeddings, + Tensor carry, + int rows, + TensorMap& env) +{ + Data& data = data_[phase]; + const cudaStream_t stream = core::Context::stream().handle(); + + Tensor concat = data.normalized_concat.slice({0, 0}, {rows, 2 * hidden_units_}); + invokeRMSNormConcat(concat, + embeddings, + spec_weights_.pre_fc_norm_embedding->weight, + spec_weights_.pre_fc_norm_embedding->norm_eps_, + spec_weights_.pre_fc_norm_embedding->zero_centered_, + carry, + spec_weights_.pre_fc_norm_hidden->weight, + spec_weights_.pre_fc_norm_hidden->norm_eps_, + spec_weights_.pre_fc_norm_hidden->zero_centered_, + stream); + + CombineResult result; + result.residual = ProjectAndGatherFc(data, concat, rows, stream, env); + return result; +} + +Tensor Qwen35MtpModel::NextCarry(const LanguageModel::DecoderOutputs& out, + int, + int, + TensorMap&, + cudaStream_t) +{ + return out.selected_hidden; +} + +LanguageModel& Qwen35MtpModel::HeadModel() +{ + return target_; +} + +Tensor Qwen35MtpModel::ProjectAndGatherFc(Data& data, + const Tensor& normalized_concat, + int rows, + cudaStream_t stream, + const TensorMap& env) +{ + const int local_hidden = hidden_units_ / model_tp_size_; + if (rows == 0) { + return data.projected_full.slice({0, 0}, {0, hidden_units_}); + } + + if (model_tp_size_ == 1) { + Tensor full = data.projected_full.slice({0, 0}, {rows, hidden_units_}); + linear_.Forward(normalized_concat, *spec_weights_.fc, full); + return full; + } + + if (use_ag2d_) { + Tensor gathered = data.projected_full.slice({0, 0}, {rows, hidden_units_}) + .view({rows, model_tp_size_, local_hidden}); + Tensor local = gathered.slice({0, model_tp_rank_, 0}, {rows, 1, local_hidden}).squeeze(1); + linear_.Forward(normalized_concat, *spec_weights_.fc, local); + comm_.d_comm->AllGather2D(local.raw_data(), + gathered.raw_data(), + hidden_units_, + local_hidden, + local_hidden, + rows, + local.dtype(), + {true, true}, + comm_.d_tp_group, + stream); + return gathered.view({rows, hidden_units_}); + } + + Tensor gathered{env.at("symm_buf").buffer().view(normalized_concat.dtype()), + {model_tp_size_, rows, local_hidden}}; + Tensor local = gathered.slice({model_tp_rank_, 0, 0}, {1, rows, local_hidden}).squeeze(0); + linear_.Forward(normalized_concat, *spec_weights_.fc, local); + comm_.d_comm->AllGather(local.raw_data(), + gathered.raw_data(), + local.size(), + local.dtype(), + comm_.d_tp_group, + stream); + + Tensor full = data.projected_full.slice({0, 0}, {rows, hidden_units_}); + invokeTransposeAxis01(static_cast(full.raw_data()), + static_cast(gathered.raw_data()), + model_tp_size_, + rows, + local_hidden, + stream); + return full; +} + +TM_REGISTER_SPECULATIVE_MODEL("mtp", Qwen35MtpModel); + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/qwen3_5_mtp/qwen3_5_mtp_model.h b/src/turbomind/models/speculative/qwen3_5_mtp/qwen3_5_mtp_model.h new file mode 100644 index 0000000000..4ff718d93c --- /dev/null +++ b/src/turbomind/models/speculative/qwen3_5_mtp/qwen3_5_mtp_model.h @@ -0,0 +1,64 @@ +// Copyright (c) OpenMMLab. All rights reserved. +#pragma once + +#include "src/turbomind/models/speculative/fixed_chain_model.h" + +#include +#include + +namespace turbomind { + +class Qwen35MtpWeight; +class TargetFinalHidden; + +class Qwen35MtpModel final: public FixedChainSpeculativeModel { +public: + explicit Qwen35MtpModel(const SpeculativeModelArgs& args); + ~Qwen35MtpModel() override; + + bool requires_successor_input_embeddings() const override + { + return true; + } + +protected: + HiddenStateTap* TapSource() override; + Tensor InitialCarry(int phase, cudaStream_t stream) override; + Tensor Embed(int phase, + const Buffer_& ids, + int rows, + const DraftContext& ctx, + TensorMap& env, + EmbedStage stage) override; + CombineResult Combine(int phase, Tensor embeddings, Tensor carry, int rows, TensorMap& env) override; + Tensor NextCarry(const LanguageModel::DecoderOutputs& out, + int step, + int phase, + TensorMap& env, + cudaStream_t stream) override; + LanguageModel& HeadModel() override; + +private: + struct Data { + Tensor normalized_concat; + Tensor projected_full; + }; + + Tensor ProjectAndGatherFc(Data& data, + const Tensor& normalized_concat, + int rows, + cudaStream_t stream, + const TensorMap& env); + + const Qwen35MtpWeight& spec_weights_; + LlamaLinear& linear_; + + const int hidden_units_; + const int model_tp_rank_; + const int model_tp_size_; + + std::unique_ptr final_hidden_; + std::vector data_; +}; + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/qwen3_5_mtp/qwen3_5_mtp_weight.cc b/src/turbomind/models/speculative/qwen3_5_mtp/qwen3_5_mtp_weight.cc new file mode 100644 index 0000000000..520fa23b40 --- /dev/null +++ b/src/turbomind/models/speculative/qwen3_5_mtp/qwen3_5_mtp_weight.cc @@ -0,0 +1,12 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#include "src/turbomind/models/speculative/qwen3_5_mtp/qwen3_5_mtp_weight.h" + +#include "src/turbomind/core/registry.h" + +namespace turbomind { + +TM_MODULE_REGISTER(Qwen35MtpWeight, core::Qwen35MtpWeightConfig); +TM_MODULE_METHODS(Qwen35MtpWeight, QWEN35_MTP_WEIGHT_CHILDREN, QWEN35_MTP_WEIGHT_PARAMS) + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/qwen3_5_mtp/qwen3_5_mtp_weight.h b/src/turbomind/models/speculative/qwen3_5_mtp/qwen3_5_mtp_weight.h new file mode 100644 index 0000000000..6f1425ce51 --- /dev/null +++ b/src/turbomind/models/speculative/qwen3_5_mtp/qwen3_5_mtp_weight.h @@ -0,0 +1,45 @@ +// Copyright (c) OpenMMLab. All rights reserved. +#pragma once + +#include "src/turbomind/core/module.h" +#include "src/turbomind/models/linear_weight.h" +#include "src/turbomind/models/norm_weight.h" + +namespace turbomind::core { + +struct Qwen35MtpWeightConfig: ModuleConfig { + Qwen35MtpWeightConfig(): ModuleConfig{"Qwen35MtpWeight"} {} + +#define QWEN35_MTP_WEIGHT_FIELDS(X) X(DataType, data_type) + + QWEN35_MTP_WEIGHT_FIELDS(TM_MEMBER) + TM_FOR_EACH(Qwen35MtpWeightConfig, QWEN35_MTP_WEIGHT_FIELDS) + +#undef QWEN35_MTP_WEIGHT_FIELDS +}; + +} // namespace turbomind::core + +namespace turbomind { + +class Qwen35MtpWeight final: public core::Module { +public: + const char* type() const override + { + return "Qwen35MtpWeight"; + } + + Qwen35MtpWeight() = default; + explicit Qwen35MtpWeight(const core::Qwen35MtpWeightConfig&) {} + +#define QWEN35_MTP_WEIGHT_CHILDREN(X) \ + X(LinearWeight, fc) \ + X(NormWeight, pre_fc_norm_embedding) \ + X(NormWeight, pre_fc_norm_hidden) + +#define QWEN35_MTP_WEIGHT_PARAMS(X) + + TM_MODULE_DECLARE(Qwen35MtpWeight, QWEN35_MTP_WEIGHT_CHILDREN, QWEN35_MTP_WEIGHT_PARAMS) +}; + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/qwen3_5_mtp/target_final_hidden.cc b/src/turbomind/models/speculative/qwen3_5_mtp/target_final_hidden.cc new file mode 100644 index 0000000000..223e9e514d --- /dev/null +++ b/src/turbomind/models/speculative/qwen3_5_mtp/target_final_hidden.cc @@ -0,0 +1,66 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#include "src/turbomind/models/speculative/qwen3_5_mtp/target_final_hidden.h" + +#include "src/turbomind/core/check.h" +#include "src/turbomind/kernels/core/math.h" +#include "src/turbomind/models/llama/context.h" + +namespace turbomind { + +TargetFinalHidden::TargetFinalHidden(const EngineParam& engine, + const Context& context, + int phases, + int target_layer_count, + int hidden_units, + DataType data_type): + HiddenStateTap{*context.is_warm_up}, + target_layer_count_{target_layer_count}, + hidden_units_{hidden_units}, + collect_{engine, context, phases, hidden_units, hidden_units, data_type} +{ +} + +void TargetFinalHidden::Begin(int phase, const std::vector& local_token_nums) +{ + active_phase_ = phase; + collect_.Begin(phase, local_token_nums); +} + +void TargetFinalHidden::Capture( + int, const Tensor&, const Tensor& local_normalized_hidden, cudaStream_t stream) +{ + const comm::OwnedTokenRows& owned = collect_.owned(active_phase_); + const int first = owned.local_begin(); + const int rows = owned.row_count(); + if (rows == 0) { + return; + } + + Tensor& captured = collect_.captured(active_phase_); + + const size_t element_bytes = byte_size(captured.dtype()); + const size_t row_bytes = size_t(hidden_units_) * element_bytes; + const auto* source = static_cast(local_normalized_hidden.raw_data()) + + size_t(first) * local_normalized_hidden.stride(0) * element_bytes; + TM_CUDA_CHECK(cudaMemcpy2DAsync(captured.raw_data(), + captured.stride(0) * element_bytes, + source, + local_normalized_hidden.stride(0) * element_bytes, + row_bytes, + rows, + cudaMemcpyDeviceToDevice, + stream)); +} + +void TargetFinalHidden::SeedWarmup(int phase, cudaStream_t stream) +{ + collect_.SeedWarmup(phase, stream); +} + +Tensor TargetFinalHidden::Gather(int phase, cudaStream_t stream) +{ + return collect_.Gather(phase, collect_.captured(phase), stream); +} + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/qwen3_5_mtp/target_final_hidden.h b/src/turbomind/models/speculative/qwen3_5_mtp/target_final_hidden.h new file mode 100644 index 0000000000..aa481e9f86 --- /dev/null +++ b/src/turbomind/models/speculative/qwen3_5_mtp/target_final_hidden.h @@ -0,0 +1,45 @@ +// Copyright (c) OpenMMLab. All rights reserved. +#pragma once + +#include "src/turbomind/core/core.h" +#include "src/turbomind/models/speculative/collect_hidden_states.h" +#include "src/turbomind/models/speculative/hidden_state_tap.h" + +namespace turbomind { + +class Context; +class EngineParam; + +/// Tap for the target's final normalized hidden: captures this rank's owned +/// rows of the last decoder layer's output and reunites them for the draft pass. +class TargetFinalHidden final: public HiddenStateTap { +public: + TargetFinalHidden(const EngineParam& engine, + const Context& context, + int phases, + int target_layer_count, + int hidden_units, + DataType data_type); + + void Begin(int phase, const std::vector& local_token_nums) override; + void SeedWarmup(int phase, cudaStream_t stream) override; + Tensor Gather(int phase, cudaStream_t stream); + + int TapOrdinal(int completed_layer_count) const override + { + return completed_layer_count == target_layer_count_ ? 0 : -1; + } + + void Capture(int, + const Tensor&, + const Tensor& local_normalized_hidden, + cudaStream_t stream) override; + +private: + int target_layer_count_; + int hidden_units_; + CollectHiddenStates collect_; + int active_phase_{}; +}; + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/registry.cc b/src/turbomind/models/speculative/registry.cc new file mode 100644 index 0000000000..0eb5cc96e7 --- /dev/null +++ b/src/turbomind/models/speculative/registry.cc @@ -0,0 +1,33 @@ +// Copyright (c) OpenMMLab. All rights reserved. + +#include "src/turbomind/models/speculative/registry.h" + +#include "src/turbomind/core/check.h" + +namespace turbomind { + +SpeculativeModelRegistry& SpeculativeModelRegistry::Instance() +{ + static SpeculativeModelRegistry registry; + return registry; +} + +void SpeculativeModelRegistry::Register(std::string name, Factory factory) +{ + factories_.emplace(std::move(name), std::move(factory)); +} + +bool SpeculativeModelRegistry::Contains(std::string_view name) const +{ + return factories_.find(name) != factories_.end(); +} + +std::unique_ptr +SpeculativeModelRegistry::Create(std::string_view name, const SpeculativeModelArgs& args) const +{ + auto it = factories_.find(name); + TM_CHECK(it != factories_.end()) << "unknown speculative method '" << name << "'"; + return it->second(args); +} + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/registry.h b/src/turbomind/models/speculative/registry.h new file mode 100644 index 0000000000..23aa041afc --- /dev/null +++ b/src/turbomind/models/speculative/registry.h @@ -0,0 +1,40 @@ +// Copyright (c) OpenMMLab. All rights reserved. +#pragma once + +#include "src/turbomind/models/speculative/speculative_model.h" + +#include +#include +#include +#include +#include + +namespace turbomind { + +class SpeculativeModelRegistry { +public: + using Factory = std::function(const SpeculativeModelArgs&)>; + + static SpeculativeModelRegistry& Instance(); + + void Register(std::string name, Factory factory); + bool Contains(std::string_view name) const; + + std::unique_ptr Create(std::string_view name, const SpeculativeModelArgs& args) const; + +private: + std::map> factories_; +}; + +} // namespace turbomind + +#define TM_REGISTER_SPECULATIVE_MODEL(name, ModelClass) \ + namespace { \ + static const bool _tm_speculative_registered_##ModelClass = [] { \ + ::turbomind::SpeculativeModelRegistry::Instance().Register( \ + name, [](const ::turbomind::SpeculativeModelArgs& args) { \ + return std::make_unique(args); \ + }); \ + return true; \ + }(); \ + } diff --git a/src/turbomind/models/speculative/speculative_model.h b/src/turbomind/models/speculative/speculative_model.h new file mode 100644 index 0000000000..d452d2e1e0 --- /dev/null +++ b/src/turbomind/models/speculative/speculative_model.h @@ -0,0 +1,66 @@ +// Copyright (c) OpenMMLab. All rights reserved. +#pragma once + +#include "src/turbomind/core/core.h" +#include "src/turbomind/engine/batch.h" +#include "src/turbomind/models/speculative/speculative_policy.h" + +namespace turbomind { + +class CacheRegistry; +struct Context; +struct EngineParam; +class HiddenStateTap; +class InputProcessor; +class LanguageModel; +class ModelWeight; + +struct DraftContext { + int batch_size; + Buffer_ target_q_offsets; + Buffer_ target_k_offsets; + Buffer_ accept_len; + Buffer_ sequence_length; + Buffer_ finished_on_entry; + int* const* request_token_ids_ptrs; + Tensor target_pre_final_residual; + Tensor embedding_storage; + Tensor head_storage; + InputProcessor* input_processor{}; +}; + +struct SpeculativeModelArgs { + CacheRegistry& registry; + const EngineParam& param; + const Context& ctx; + LanguageModel& target; + const ModelWeight& draft_weights; + int phases; +}; + +class SpeculativeModel { +public: + virtual ~SpeculativeModel() = default; + + virtual const SpeculativePolicy& policy() const = 0; + + /// The draft model's batch-op lifecycle: kAdd, kSetup, and kPrepare. + virtual void Run(BatchOp op, int phase, TensorMap& env) = 0; + + /// The one target-pass hook: the tap the target decoder captures through, + /// armed by the executor before the target pass. Null when the policy + /// taps no layer. + virtual HiddenStateTap* Tap(int phase) + { + return nullptr; + } + + virtual bool requires_successor_input_embeddings() const + { + return false; + } + + virtual void RunDraft(int phase, const DraftContext& ctx, TensorMap& env) = 0; +}; + +} // namespace turbomind diff --git a/src/turbomind/models/speculative/speculative_policy.h b/src/turbomind/models/speculative/speculative_policy.h new file mode 100644 index 0000000000..d78c1403ed --- /dev/null +++ b/src/turbomind/models/speculative/speculative_policy.h @@ -0,0 +1,58 @@ +// Copyright (c) OpenMMLab. All rights reserved. +#pragma once + +#include + +namespace turbomind { + +/// What one speculative round costs in rows and cache. +struct RoundExtent { + int query_rows; // target query rows the round submits + int private_tail; // cache rows written past key capacity, never committed +}; + +/// Extra room the final prompt forward needs so it can write the first proposals. +struct BootstrapExtent { + int cache_write_end; + int min_session_len; +}; + +/// Fixed-width methods ignore this. It carries what the scheduler already knows +/// about the row being planned, so a width that depends on position within a +/// sequence needs no signature change. Adapting to a request's own acceptance +/// history additionally needs request identity and a feedback call; see Deferred. +struct RoundRequest { + int query_begin; // where this round's first query row lands + int prompt_len; +}; + +/// The scheduling and verification contract of a speculative method. Consumed by +/// Scheduler and Engine, neither of which knows the algorithm. +class SpeculativePolicy { +public: + virtual ~SpeculativePolicy() = default; + + /// Upper bound on proposals per round. Sizes verification buffers and the + /// selected-span stride; not the count actually proposed in a given round. + virtual int max_proposals() const = 0; + + /// Extra token-row columns past session_len that proposal writes need. Sizes + /// Generation's token row; distinct from RoundExtent::private_tail, which sizes + /// the KV cache. + virtual int token_row_tail() const = 0; + + virtual RoundExtent Extent(const RoundRequest& req) const = 0; + + /// Nullopt when the method does not bootstrap from the prompt pass. + virtual std::optional Bootstrap(int prompt_len) const = 0; + + /// True when the final prompt forward must have a token row already allocated + /// so it can write the first proposals into it. Read by Generation, which + /// otherwise allocates a row only for generating requests. + bool needs_prompt_token_row(int prompt_len) const + { + return Bootstrap(prompt_len).has_value(); + } +}; + +} // namespace turbomind diff --git a/src/turbomind/models/vision_model.cc b/src/turbomind/models/vision_model.cc index 4f2e10dcf4..db6a6d80a8 100644 --- a/src/turbomind/models/vision_model.cc +++ b/src/turbomind/models/vision_model.cc @@ -15,13 +15,16 @@ namespace turbomind { std::unique_ptr CreateVisionModel(const VisionModelWeight& weights, // const EngineParam& engine, const Context& ctx, - int phases) + int phases, + bool successor_embeddings) { if (std::string_view{weights.type()} == "QwenVitWeight") { - return std::make_unique(engine, ctx, static_cast(weights), phases); + return std::make_unique( + engine, ctx, static_cast(weights), phases, successor_embeddings); } if (std::string_view{weights.type()} == "InternVitWeight") { - return std::make_unique(engine, ctx, static_cast(weights), phases); + return std::make_unique( + engine, ctx, static_cast(weights), phases, successor_embeddings); } TM_LOG_FATAL("Unsupported vision model weight type: {}", weights.type()); diff --git a/src/turbomind/models/vision_model.h b/src/turbomind/models/vision_model.h index a1c2edc9aa..cd1b9f1464 100644 --- a/src/turbomind/models/vision_model.h +++ b/src/turbomind/models/vision_model.h @@ -28,7 +28,8 @@ class VisionModel { public: virtual ~VisionModel() = default; - /// Phase entry point. Called from ``ModelExecutor::Run`` *before* + /// Phase entry point. Called from the batch-operation fanouts (the + /// engine's host-op functions and the executor's device steps) *before* /// the language model. Subclasses dispatch on ``op``. virtual void Run(BatchOp op, int phase, TensorMap& env) = 0; }; @@ -39,19 +40,25 @@ struct MultiModalData { std::array grid_thw; // qwen3 }; +struct EmbeddingPatch { + int row_count; + int source_row; + int destination_row; +}; + struct MultiModalEmbeddingData { - Tensor data; - std::vector> image_embeds_coords; - std::vector> input_embeds_coords; + Tensor data; + std::vector target_patches; + std::vector successor_patches; MultiModalEmbeddingData() = default; - explicit MultiModalEmbeddingData(Tensor data, - std::vector> image_embeds_coords, - std::vector> input_embeds_coords): + explicit MultiModalEmbeddingData(Tensor data, + std::vector target_patches, + std::vector successor_patches): data{std::move(data)}, - image_embeds_coords{std::move(image_embeds_coords)}, - input_embeds_coords{std::move(input_embeds_coords)} + target_patches{std::move(target_patches)}, + successor_patches{std::move(successor_patches)} { } @@ -80,6 +87,7 @@ struct MultiModalEmbeddingData { std::unique_ptr CreateVisionModel(const VisionModelWeight& weights, // const EngineParam& engine, const Context& ctx, - int phases); + int phases, + bool successor_embeddings); } // namespace turbomind diff --git a/src/turbomind/python/CMakeLists.txt b/src/turbomind/python/CMakeLists.txt index 6f2a6606d6..74caad11b3 100644 --- a/src/turbomind/python/CMakeLists.txt +++ b/src/turbomind/python/CMakeLists.txt @@ -17,8 +17,21 @@ pybind11_add_module(${PROJECT_NAME} linear_bind.cpp xgrammar_bind.cpp ../kernels/linear_attn/python_bind.cpp - ../kernels/gemm/moe_gate_python_bind.cpp) -target_link_libraries(${PROJECT_NAME} PRIVATE turbomind xgrammar) + ../kernels/gemm/moe_gate_python_bind.cpp + ../kernels/speculative_sampling_python_bind.cpp + ../kernels/draft_carry_python_bind.cpp + ../models/speculative/eagle3/target_hidden_projection_python_bind.cpp + ../kernels/speculative_sequence_python_bind.cpp + ../kernels/attention/verification/python_bind.cpp) +target_link_libraries(${PROJECT_NAME} PRIVATE + turbomind + speculative_sampling_kernels + draft_carry_kernels + target_hidden_projection_kernels + speculative_sequence_kernels + verification_attention + attention + xgrammar) string(REPLACE "." ";" _cuda_version ${CMAKE_CUDA_COMPILER_VERSION}) list(GET _cuda_version 0 CUDA_MAJOR) diff --git a/src/turbomind/python/attention_component_bindings.h b/src/turbomind/python/attention_component_bindings.h new file mode 100644 index 0000000000..e4aae754d5 --- /dev/null +++ b/src/turbomind/python/attention_component_bindings.h @@ -0,0 +1,9 @@ +#pragma once + +#include + +namespace turbomind::python { + +void BindVerificationAttention(pybind11::module_& module); + +} // namespace turbomind::python diff --git a/src/turbomind/python/bind.cpp b/src/turbomind/python/bind.cpp index 1510106bba..1dc94f3f6a 100644 --- a/src/turbomind/python/bind.cpp +++ b/src/turbomind/python/bind.cpp @@ -42,8 +42,12 @@ #include "src/turbomind/models/qwenvit/qwenvit_block_weight.h" #include "src/turbomind/models/qwenvit/qwenvit_input.h" #include "src/turbomind/models/qwenvit/qwenvit_weight.h" +#include "src/turbomind/models/speculative/eagle3/eagle3_weight.h" +#include "src/turbomind/models/speculative/qwen3_5_mtp/qwen3_5_mtp_weight.h" #include "src/turbomind/models/vision_model_weight.h" #include "src/turbomind/python/dlpack.h" +#include "src/turbomind/python/attention_component_bindings.h" +#include "src/turbomind/python/eagle3_component_bindings.h" #include "src/turbomind/turbomind.h" #include "src/turbomind/utils/cuda_utils.h" #include "src/turbomind/utils/metrics.h" @@ -306,6 +310,80 @@ std::shared_ptr FromDLPack(const py::object& object) return std::make_shared(std::move(owner), std::move(layout), dtype, device); } +ft::core::ssize_t TorchStorageCapacityElements(py::handle source, ft::DataType dtype, ft::core::ssize_t fallback) +{ + if (!source || !py::hasattr(source, "untyped_storage") || !py::hasattr(source, "storage_offset")) { + return fallback; + } + + try { + const auto elem_bytes = ft::byte_size(dtype); + if (elem_bytes <= 0) { + return fallback; + } + auto storage = source.attr("untyped_storage")(); + auto storage_bytes = py::cast(storage.attr("nbytes")()); + auto offset = py::cast(source.attr("storage_offset")()); + if (storage_bytes < 0 || offset < 0) { + return fallback; + } + const auto offset_bytes = offset * elem_bytes; + if (offset_bytes < 0 || offset_bytes > storage_bytes) { + return fallback; + } + const auto capacity = (storage_bytes - offset_bytes) / elem_bytes; + return capacity > fallback ? capacity : fallback; + } + catch (py::error_already_set& e) { + e.restore(); + PyErr_Clear(); + return fallback; + } +} + +// Like FromDLPack, but sizes the buffer to the source storage's capacity so +// strided views (e.g. transposed destinations) can be written past cosize(). +std::shared_ptr FromDLPackWithStrides(const py::object& object) +{ + py::capsule capsule = object.attr("__dlpack__")(); + auto* managed = static_cast(PyCapsule_GetPointer(capsule.ptr(), kDlTensorCapsuleName)); + auto& dl_tensor = managed->dl_tensor; + + const ft::core::Device device{getMemoryType(dl_tensor.device), dl_tensor.device.device_id}; + const auto dtype = getDataType(dl_tensor.dtype); + assert(dl_tensor.ndim > 0); + std::vector shape(dl_tensor.shape, dl_tensor.shape + dl_tensor.ndim); + + // Compute row-major strides if DLPack strides are NULL (contiguous tensor) + std::vector strides; + if (dl_tensor.strides) { + strides.assign(dl_tensor.strides, dl_tensor.strides + dl_tensor.ndim); + } + else { + strides.resize(dl_tensor.ndim); + ft::core::ssize_t value = 1; + for (int i = dl_tensor.ndim - 1; i >= 0; --i) { + strides[i] = value; + value *= shape[i]; + } + } + + ft::core::Layout layout{std::move(shape), std::move(strides)}; + auto* data = static_cast(dl_tensor.data) + dl_tensor.byte_offset; + + std::shared_ptr owner{data, [managed](void*) { + if (managed->deleter) { + managed->deleter(managed); + } + }}; + capsule.set_name("used_dltensor"); + + const auto capacity = + layout.is_contiguous() ? layout.cosize() : TorchStorageCapacityElements(object, dtype, layout.cosize()); + auto buffer = ft::core::Buffer{std::move(owner), capacity, dtype, device}; + return std::make_shared(std::move(buffer), std::move(layout), Tensor::PreserveBufferCapacity{}); +} + static void safe_memcpy(void* dst, const void* src, size_t size) { cudaPointerAttributes dat{}; @@ -548,7 +626,11 @@ PYBIND11_MODULE(_turbomind, m) py::class_>(m, "RequestState") .def_readonly("status", &ft::RequestState::status) - .def_readonly("seq_len", &ft::RequestState::seq_len); + .def_readonly("seq_len", &ft::RequestState::seq_len) + .def_readonly("num_drafts", &ft::RequestState::num_drafts) + .def_readonly("num_draft_tokens", &ft::RequestState::num_draft_tokens) + .def_readonly("num_accepted_tokens", &ft::RequestState::num_accepted_tokens) + .def_readonly("num_accepted_tokens_per_pos", &ft::RequestState::num_accepted_tokens_per_pos); py::class_>(m, "AtomicRequestState") .def("consume", [](ft::AtomicRequestState& s) { return s.exchange(nullptr); }); @@ -624,6 +706,8 @@ PYBIND11_MODULE(_turbomind, m) bind_config(m, "ModuleListConfig"); bind_config(m, "NormConfig"); bind_config(m, "DecoderLayerConfig"); + bind_config(m, "Eagle3WeightConfig"); + bind_config(m, "Qwen35MtpWeightConfig"); bind_config(m, "ModelWeightConfig"); bind_config(m, "LayerNormConfig"); bind_config(m, "QwenVitConfig"); @@ -696,6 +780,7 @@ PYBIND11_MODULE(_turbomind, m) "dtype"_a, "shape"_a); m.def("from_dlpack", &FromDLPack, "tensor"_a); + m.def("from_dlpack_with_strides", &FromDLPackWithStrides, "tensor"_a); m.def( "generic_copy_on_stream", [](std::shared_ptr src, std::shared_ptr dst, std::uintptr_t stream_ptr) { @@ -970,4 +1055,9 @@ PYBIND11_MODULE(_turbomind, m) turbomind::linear_attn::delta_rule::bind_delta_rule(m); turbomind::python_linear::bind_linear(m); turbomind::bind_moe_gate_v2(m); + turbomind::python::BindSpeculativeSampling(m); + turbomind::python::BindDraftCarry(m); + turbomind::python::BindTargetHiddenProjection(m); + turbomind::python::BindSpeculativeSequence(m); + turbomind::python::BindVerificationAttention(m); } diff --git a/src/turbomind/python/eagle3_component_bindings.h b/src/turbomind/python/eagle3_component_bindings.h new file mode 100644 index 0000000000..8254480095 --- /dev/null +++ b/src/turbomind/python/eagle3_component_bindings.h @@ -0,0 +1,15 @@ +#pragma once + +#include + +namespace turbomind::python { + +void BindSpeculativeSampling(pybind11::module_& module); + +void BindDraftCarry(pybind11::module_& module); + +void BindTargetHiddenProjection(pybind11::module_& module); + +void BindSpeculativeSequence(pybind11::module_& module); + +} // namespace turbomind::python diff --git a/src/turbomind/python/eagle3_dlpack_internal.h b/src/turbomind/python/eagle3_dlpack_internal.h new file mode 100644 index 0000000000..2f46d45aab --- /dev/null +++ b/src/turbomind/python/eagle3_dlpack_internal.h @@ -0,0 +1,147 @@ +#pragma once + +#include +#include +#include +#include +#include + +#include + +#include "src/turbomind/core/data_type.h" +#include "src/turbomind/core/tensor.h" +#include "src/turbomind/python/dlpack.h" + +namespace py = pybind11; +namespace ft = turbomind; + +namespace turbomind::python::detail { + +using ft::core::Tensor; + +inline constexpr char kDlTensorCapsuleName[] = "dltensor"; + +inline ft::DataType getDataType(DLDataType source) +{ + using ft::data_type_v; + switch (source.code) { + case DLDataTypeCode::kDLUInt: + switch (source.bits) { + case 8: + return data_type_v; + case 16: + return data_type_v; + case 32: + return data_type_v; + case 64: + return data_type_v; + } + break; + case DLDataTypeCode::kDLInt: + switch (source.bits) { + case 8: + return data_type_v; + case 16: + return data_type_v; + case 32: + return data_type_v; + case 64: + return data_type_v; + } + break; + case DLDataTypeCode::kDLFloat: + switch (source.bits) { + case 16: + return data_type_v; + case 32: + return data_type_v; + case 64: + return data_type_v; + } + break; + case DLDataTypeCode::kDLBfloat: + if (source.bits == 16) { + return data_type_v; + } + break; + case DLDataTypeCode::kDLBool: + if (source.bits == 8) { + return data_type_v; + } + break; + } + __builtin_unreachable(); +} + +inline ft::core::Device getDevice(DLDevice source) +{ + switch (source.device_type) { + case DLDeviceType::kDLCUDA: + return {ft::DeviceType::kDEVICE, source.device_id}; + case DLDeviceType::kDLCUDAHost: + return {ft::DeviceType::kCPUpinned, -1}; + case DLDeviceType::kDLCPU: + return {ft::DeviceType::kCPU, -1}; + default: + __builtin_unreachable(); + } +} + +inline Tensor ConsumeDLPackWithStrides(py::handle object, uintptr_t consumer_stream) +{ + const uintptr_t dlpack_stream = consumer_stream == 0 ? 1 : consumer_stream; + + py::capsule capsule = object.attr("__dlpack__")(py::arg("stream") = py::int_(dlpack_stream)); + auto* managed = static_cast(PyCapsule_GetPointer(capsule.ptr(), kDlTensorCapsuleName)); + auto& source = managed->dl_tensor; + + using index_t = ft::core::ssize_t; + std::vector shape(source.ndim); + std::vector stride(source.ndim); + + for (int i = 0; i < source.ndim; ++i) { + shape[i] = static_cast(source.shape[i]); + } + + if (source.strides) { + for (int i = 0; i < source.ndim; ++i) { + stride[i] = static_cast(source.strides[i]); + } + } + else { + stride.back() = 1; + if (source.ndim == 2) { + stride.front() = shape.back(); + } + } + + index_t logical_size = shape[0]; + if (source.ndim == 2) { + logical_size *= shape[1]; + } + if (logical_size == 0) { + stride.back() = 1; + if (source.ndim == 2) { + stride.front() = shape.back(); + } + } + + void* data = source.data; + if (source.byte_offset != 0) { + data = reinterpret_cast(reinterpret_cast(data) + static_cast(source.byte_offset)); + } + + capsule.set_name("used_dltensor"); + std::shared_ptr owner{data, [managed](void*) { + if (managed->deleter) { + managed->deleter(managed); + } + }}; + + return Tensor{std::move(owner), + ft::core::Layout{std::move(shape), std::move(stride)}, + getDataType(source.dtype), + getDevice(source.device)}; +} + +} // namespace turbomind::python::detail diff --git a/src/turbomind/turbomind.cc b/src/turbomind/turbomind.cc index c73d4b3ce4..5f959f6f02 100644 --- a/src/turbomind/turbomind.cc +++ b/src/turbomind/turbomind.cc @@ -1,7 +1,6 @@ // Copyright (c) OpenMMLab. All rights reserved. #include -#include #include #include "src/turbomind/turbomind.h" @@ -15,7 +14,6 @@ #include "src/turbomind/engine/cache_registry.h" #include "src/turbomind/engine/engine.h" #include "src/turbomind/engine/gateway.h" -#include "src/turbomind/engine/model_executor.h" #include "src/turbomind/engine/model_request.h" #include "src/turbomind/models/language_model.h" @@ -23,6 +21,7 @@ #include "src/turbomind/models/llama/llama_params.h" #include "src/turbomind/models/model_root.h" #include "src/turbomind/models/model_weight.h" +#include "src/turbomind/models/speculative/registry.h" #include "src/turbomind/models/vision_model.h" #include "src/turbomind/kernels/gemm/tuner/params.h" @@ -79,7 +78,8 @@ struct TurboMind::Impl { data_type_, engine_param_.session_len, weights_[0]->text_model_ptr()->vocab_size, - weights_[0]->text_model_ptr()->hidden_units); + weights_[0]->text_model_ptr()->hidden_units, + engine_param_.spec_num_draft_tokens); } core::Module* CreateRoot(int index) @@ -323,53 +323,50 @@ void TurboMind::Impl::CreateEngine(int index) ctx.comm.h_comm->Sync(); - const double cache_ratio = param.cache_max_block_count; - TM_CHECK_GT(cache_ratio, 0.) << "object-cache path expects 0 < cache_max_block_count < 1"; - TM_CHECK_LT(cache_ratio, 1.) << "object-cache path no longer accepts cache_max_block_count as a block count"; - - size_t free_bytes{}, total_bytes{}; - TM_CUDA_CHECK(cudaMemGetInfo(&free_bytes, &total_bytes)); - free_bytes = AllReduce(ctx.comm.h_tp_group, free_bytes, comm::RedOp::kMin); + CacheRegistry cache_registry; + cache_registry.set_checkpoint_min_interval(param.cache_checkpoint_interval); - const size_t cache_bytes = static_cast(static_cast(free_bytes) * cache_ratio); - TM_CHECK_GT(cache_bytes, size_t{0}); - TM_CHECK_LE(cache_bytes, static_cast(std::numeric_limits::max())); + const ModelWeight* draft_weights = weights_[index]->draft_model_ptr(); + const ModelWeight& target_weights = *weights_[index]->text_model_ptr(); - TM_LOG_INFO("Object cache budget: {:.2f} MB from free {:.2f} MB and ratio {:.3f}", - cache_bytes / (1024. * 1024.), - free_bytes / (1024. * 1024.), - cache_ratio); + const size_t prefix_before = cache_registry.prefix().accumulation_bytes(); + auto model = std::make_unique(cache_registry, param, ctx, target_weights, phases_); + const size_t target_end = cache_registry.prefix().accumulation_bytes(); - Buffer cache_region{static_cast(cache_bytes), data_type_v, core::Context::device_alloc()}; - ObjectAllocator alloc{std::move(cache_region)}; - CacheRegistry cache_registry; - cache_registry.set_checkpoint_min_interval(param.cache_checkpoint_interval); + std::unique_ptr spec_model; + if (!param.spec_method.empty()) { + auto& registry = SpeculativeModelRegistry::Instance(); + TM_CHECK(registry.Contains(param.spec_method)) << "unknown speculative method '" << param.spec_method << "'"; + TM_CHECK(draft_weights) << "speculative method '" << param.spec_method << "' has no draft weight tree"; + spec_model = registry.Create( + param.spec_method, {cache_registry, param, ctx, *model, *draft_weights, phases_}); + } + const size_t combined_end = cache_registry.prefix().accumulation_bytes(); + const bool successor_embeddings = + spec_model && spec_model->requires_successor_input_embeddings(); - // create model - LanguageModel model{cache_registry, param, ctx, *weights_[index]->text_model_ptr(), phases_}; + TM_LOG_INFO("Target KV block bytes: {}", target_end - prefix_before); + TM_LOG_INFO("Draft KV block bytes: {}", combined_end - target_end); + TM_LOG_INFO("Total KV block bytes: {}", combined_end - prefix_before); // create vision model for VLM checkpoints; null for text-only (no vision sub-tree attached) std::unique_ptr vision_model; if (auto* vw = weights_[index]->vision_model_ptr()) { - vision_model = CreateVisionModel(*vw, param, ctx, phases_); + vision_model = CreateVisionModel(*vw, param, ctx, phases_, successor_embeddings); } - cache_registry.RegisterObjectIds(alloc); - // create engine engines_[index] = Engine{param, - std::move(alloc), std::move(cache_registry), std::move(model), std::move(vision_model), + std::move(spec_model), ctx, *gateway_, engine_param_.devices[index], queue_id_[index], phases_}; - core::Context::stream().Sync(); - ctx.comm.h_comm->Sync(); engines_[index].Start(); diff --git a/src/turbomind/utils/metrics.h b/src/turbomind/utils/metrics.h index c52c747672..4388020ee1 100644 --- a/src/turbomind/utils/metrics.h +++ b/src/turbomind/utils/metrics.h @@ -3,7 +3,9 @@ #include #include #include +#include #include +#include namespace turbomind { @@ -20,10 +22,23 @@ struct ScheduleMetrics { }; struct RequestMetrics { + explicit RequestMetrics(int speculative_tokens = 0): + num_accepted_tokens_per_pos(static_cast(speculative_tokens)) + { + } + std::atomic enqueue_time{}; // when a request is enqued std::atomic scheduled_time{}; // when a request is scheduled for inference std::atomic cached_tokens{}; // prompt tokens skipped at first admission + std::mutex spec_mutex; + + int64_t num_drafts{}; + int64_t num_draft_tokens{}; + int64_t num_accepted_tokens{}; + + std::vector num_accepted_tokens_per_pos; + static int64_t timestamp() { // Get current timestamp in microseconds since Unix epoch diff --git a/src/turbomind/utils/nvtx_utils.cc b/src/turbomind/utils/nvtx_utils.cc index 64d3d49fc1..bd4e6e50bd 100644 --- a/src/turbomind/utils/nvtx_utils.cc +++ b/src/turbomind/utils/nvtx_utils.cc @@ -66,16 +66,21 @@ bool isEnableNvtx() return is_enable_ft_nvtx; } -void ftNvtxRangePush(std::string name) +void ftNvtxRangePush(std::string_view name) { #ifdef USE_NVTX - nvtxStringHandle_t nameId = nvtxDomainRegisterStringA(NULL, (getScope() + name).c_str()); - nvtxEventAttributes_t eventAttrib = {0}; - eventAttrib.messageType = NVTX_MESSAGE_TYPE_REGISTERED; - eventAttrib.message.registered = nameId; - eventAttrib.payloadType = NVTX_PAYLOAD_TYPE_INT32; - eventAttrib.payload.iValue = getDeviceDomain(); - nvtxRangePushEx(&eventAttrib); + auto registered_name = getScope(); + registered_name.append(name.data(), name.size()); + + nvtxStringHandle_t name_id = nvtxDomainRegisterStringA(nullptr, registered_name.c_str()); + nvtxEventAttributes_t attributes = {0}; + attributes.messageType = NVTX_MESSAGE_TYPE_REGISTERED; + attributes.message.registered = name_id; + attributes.payloadType = NVTX_PAYLOAD_TYPE_INT32; + attributes.payload.iValue = getDeviceDomain(); + nvtxRangePushEx(&attributes); +#else + (void)name; #endif } diff --git a/src/turbomind/utils/nvtx_utils.h b/src/turbomind/utils/nvtx_utils.h index e000c2157b..5836409a7d 100644 --- a/src/turbomind/utils/nvtx_utils.h +++ b/src/turbomind/utils/nvtx_utils.h @@ -16,6 +16,9 @@ #pragma once +#include +#include + namespace ft_nvtx { static std::string scope; std::string getScope(); @@ -30,10 +33,36 @@ bool isEnableNvtx(); static bool has_read_nvtx_env = false; static bool is_enable_ft_nvtx = false; -void ftNvtxRangePush(std::string name); +void ftNvtxRangePush(std::string_view name); void ftNvtxRangePop(); } // namespace ft_nvtx +namespace turbomind { + +struct NvtxScope { + explicit NvtxScope(std::string_view name): active_{ft_nvtx::isEnableNvtx()} + { + if (active_) { + ft_nvtx::ftNvtxRangePush(name); + } + } + + NvtxScope(const NvtxScope&) = delete; + NvtxScope& operator=(const NvtxScope&) = delete; + + ~NvtxScope() + { + if (active_) { + ft_nvtx::ftNvtxRangePop(); + } + } + +private: + bool active_; +}; + +} // namespace turbomind + #define PUSH_RANGE(name) \ { \ if (ft_nvtx::isEnableNvtx()) { \ diff --git a/tests/turbomind/attention/__init__.py b/tests/turbomind/attention/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/tests/turbomind/attention/test_verification_attention.py b/tests/turbomind/attention/test_verification_attention.py new file mode 100644 index 0000000000..def78ed965 --- /dev/null +++ b/tests/turbomind/attention/test_verification_attention.py @@ -0,0 +1,380 @@ +import math + +import pytest +import torch + +import _turbomind as _tm + +from .verification_attention import run_verification_attention + + +def _supports_verification_attention(): + if not torch.cuda.is_available(): + return False + properties = torch.cuda.get_device_properties(torch.cuda.current_device()) + optin_shared_memory = getattr(properties, 'shared_memory_per_block_optin', + properties.shared_memory_per_block) + return (hasattr(_tm, 'verification_attention') + and properties.major == 9 + and optin_shared_memory >= 112 * 1024) + + +pytestmark = pytest.mark.skipif( + not _supports_verification_attention(), + reason='verification attention requires SM90 and 112 KiB opt-in shared memory') + + +def _rotate(x, + position, + *, + dim, + base, + factor, + mrope_mode=0, + section=(0, 0, 0)): + if dim == 0: + return x + pairs = torch.arange(dim // 2, device=x.device) + frequency = factor**-1 * base**(-(2 * pairs.float()) / dim) + if mrope_mode == 0: + timestep = position[..., None].float() + elif mrope_mode == 1: + limits = torch.tensor(section, device=x.device).cumsum(0) + axis = torch.where(pairs < limits[0], 0, + torch.where(pairs < limits[1], 1, 2)) + timestep = position[..., axis].float() + else: + cycle = pairs // 3 + axis = pairs % 3 + use_axis = torch.where((axis == 1) & (cycle < section[1]), 1, + torch.where((axis == 2) & (cycle < section[2]), + 2, 0)) + timestep = position[..., use_axis].float() + angle = timestep * frequency + source = x[..., :dim].float().reshape(*x.shape[:-1], dim // 2, 2) + first = torch.cos(angle) * source[..., 0] - torch.sin(angle) * source[..., + 1] + second = torch.cos(angle) * source[..., 1] + torch.sin(angle) * source[..., + 0] + rotated = torch.stack((first, second), dim=-1).flatten(-2).to(x.dtype) + return torch.cat((rotated, x[..., dim:]), dim=-1) + + +def _positions(request, length, *, mrope_mode, position_ids, position_delta, + mrope_length, device): + scalar = torch.arange(length, device=device) + if mrope_mode == 0: + return scalar + explicit = scalar < int(mrope_length[request]) + fallback = (scalar + int(position_delta[request]))[:, None].expand(-1, 3) + return torch.where(explicit[:, None], position_ids[request, :length], + fallback) + + +def _reference(case, prefix_k, prefix_v, packed_qkv, q_bias, + mrope_position_ids, mrope_position_delta, mrope_length): + outputs = [] + p = case['p'] + hq = case['hq'] + hk = case['hk'] + d = case['d'] + qkv_heads = hq + 2 * hk + prefix_begin = 0 + query_begin = 0 + for request, history in enumerate(case['histories']): + prefix_end = prefix_begin + history + query_end = query_begin + p + prefix_positions = _positions(request, + history, + mrope_mode=case['mrope_mode'], + position_ids=mrope_position_ids, + position_delta=mrope_position_delta, + mrope_length=mrope_length, + device=packed_qkv.device) + all_positions = _positions(request, + history + p, + mrope_mode=case['mrope_mode'], + position_ids=mrope_position_ids, + position_delta=mrope_position_delta, + mrope_length=mrope_length, + device=packed_qkv.device) + query_positions = all_positions[history:] + + q = packed_qkv[query_begin:query_end, :hq].clone() + if q_bias.numel(): + q = (q + q_bias).to(q.dtype) + q = _rotate(q, + query_positions[:, None], + dim=case['rope_dim'], + base=case['rope_base'], + factor=case['rope_factor'], + mrope_mode=case['mrope_mode'], + section=case['mrope_section']) + tail_k = packed_qkv[query_begin:query_end, hq:hq + hk] + k = torch.cat((prefix_k[prefix_begin:prefix_end], tail_k), dim=0) + k = _rotate(k, + all_positions[:, None], + dim=case['rope_dim'], + base=case['rope_base'], + factor=case['rope_factor'], + mrope_mode=case['mrope_mode'], + section=case['mrope_section']) + tail_v = packed_qkv[query_begin:query_end, hq + hk:qkv_heads] + v = torch.cat((prefix_v[prefix_begin:prefix_end], tail_v), dim=0) + + q_group = q.reshape(p, hk, hq // hk, d).float() + score = torch.einsum('thgd,shd->thgs', q_group, + k.float()) / math.sqrt(d) + keys = torch.arange(history + p, device=q.device) + queries = history + torch.arange(p, device=q.device) + valid = keys[None, :] <= queries[:, None] + if case['window']: + valid &= keys[None, :] >= queries[:, None] - case['window'] + 1 + score = score.masked_fill(~valid[:, None, None, :], float('-inf')) + probability = score.softmax(-1) + out = torch.einsum('thgs,shd->thgd', probability, v.float()) + outputs.append(out.reshape(p, hq, d).to(packed_qkv.dtype)) + prefix_begin = prefix_end + query_begin = query_end + return torch.cat(outputs) + + +def _make_case(*, + dtype=torch.bfloat16, + d=128, + p=8, + hq=4, + hk=1, + histories=(63, 257, 1025), + window=0, + bias=False, + rope_type=0, + rope_dim=0, + mrope_mode=0, + mrope_section=(0, 0, 0), + requested_splits=128): + return dict(dtype=dtype, + d=d, + p=p, + hq=hq, + hk=hk, + histories=histories, + window=window, + bias=bias, + rope_type=rope_type, + rope_dim=rope_dim, + rope_base=1_000_000., + rope_factor=1., + mrope_mode=mrope_mode, + mrope_section=mrope_section, + requested_splits=requested_splits) + + +def _run(case, finished=None): + torch.manual_seed(91) + device = torch.device('cuda') + dtype = case['dtype'] + b = len(case['histories']) + d, p, hq, hk = case['d'], case['p'], case['hq'], case['hk'] + sum_history = sum(case['histories']) + sum_query = b * p + scale = 0.2 + prefix_k = torch.randn(sum_history, hk, d, device=device, + dtype=dtype) * scale + prefix_v = torch.randn_like(prefix_k) * scale + packed_qkv = torch.randn( + sum_query, hq + 2 * hk, d, device=device, dtype=dtype) * scale + q_bias = (torch.randn(hq, d, device=device, dtype=dtype) * scale + if case['bias'] else torch.empty(0, device=device, dtype=dtype)) + + prefix_offsets = torch.tensor( + [0] + list(torch.tensor(case['histories']).cumsum(0).tolist()), + device=device, + dtype=torch.int32) + q_offsets = torch.arange(0, (b + 1) * p, + p, + device=device, + dtype=torch.int32) + key_lengths = [history + p for history in case['histories']] + k_offsets = torch.tensor( + [0] + list(torch.tensor(key_lengths).cumsum(0).tolist()), + device=device, + dtype=torch.int32) + if finished is None: + finished = torch.zeros(b, device=device, dtype=torch.bool) + + block_len = 64 + page_counts = [(length + block_len - 1) // block_len + for length in key_lengths] + block_ptr_offsets = torch.tensor( + [0] + list(torch.tensor(page_counts).cumsum(0).tolist()), + device=device, + dtype=torch.int32) + page_count = sum(page_counts) + page_bytes = 2 * hk * block_len * d * torch.tensor( + [], dtype=dtype).element_size() + cache = torch.full((page_count, page_bytes), + 0xA5, + device=device, + dtype=torch.uint8) + permutation = torch.randperm(page_count, device='cpu').tolist() + base = cache.data_ptr() + block_ptrs = torch.tensor( + [base + physical * page_bytes for physical in permutation], + device=device, + dtype=torch.int64) + + max_key_length = max(key_lengths) + mrope_position_ids = torch.empty((b, max_key_length, 3), + device=device, + dtype=torch.int32) + for request, length in enumerate(key_lengths): + pos = torch.arange(length, device=device, dtype=torch.int32) + mrope_position_ids[request, :length, 0] = pos + mrope_position_ids[request, :length, 1] = pos * 2 + 1 + mrope_position_ids[request, :length, 2] = pos * 3 + 2 + mrope_position_delta = torch.tensor([3, 5, 7], + device=device, + dtype=torch.int32) + mrope_length = torch.tensor( + [key_lengths[0], case['histories'][1] + p // 2, key_lengths[2]], + device=device, + dtype=torch.int32) + + output = torch.full((sum_query, hq, d), + float('nan'), + device=device, + dtype=dtype) + partial_capacity = 4096 + partial_o = torch.full((partial_capacity, hq, d), + float('nan'), + device=device) + partial_ml = torch.full((partial_capacity, hq, 2), + float('nan'), + device=device) + split_count = run_verification_attention( + prefix_k=prefix_k, + prefix_v=prefix_v, + prefix_offsets=prefix_offsets, + packed_qkv=packed_qkv, + q_bias=q_bias, + output=output, + cache_storage=cache, + block_ptrs=block_ptrs, + block_ptr_offsets=block_ptr_offsets, + q_offsets=q_offsets, + k_offsets=k_offsets, + finished=finished, + partial_o=partial_o, + partial_ml=partial_ml, + query_head_count=hq, + kv_head_count=hk, + head_dim=d, + block_len=block_len, + max_query_length=p, + max_key_length=max_key_length, + window_size=case['window'], + requested_max_split_count=case['requested_splits'], + rope_type=case['rope_type'], + rope_dim=case['rope_dim'], + rope_base=case['rope_base'], + rope_factor=case['rope_factor'], + mrope_mode=case['mrope_mode'], + mrope_section=case['mrope_section'], + mrope_position_ids=mrope_position_ids, + mrope_position_delta=mrope_position_delta, + mrope_length=mrope_length) + torch.cuda.synchronize() + reference = _reference(case, prefix_k, prefix_v, packed_qkv, q_bias, + mrope_position_ids, mrope_position_delta, + mrope_length) + return output, reference, partial_o, partial_ml, split_count + + +@pytest.mark.parametrize('case', [ + _make_case(d=128, p=8, hq=4, bias=True, requested_splits=8), + _make_case(d=256, p=8, hq=6), + _make_case(d=256, p=16, hq=8, histories=(65, 511, 1025)), + _make_case(dtype=torch.float16, d=256, p=4, hq=16, window=129), +]) +def test_verification_attention(case): + output, reference, _, _, split_count = _run(case) + if case['d'] == 128: + sm_count = torch.cuda.get_device_properties( + torch.cuda.current_device()).multi_processor_count + expected_split_count = min(8, (2 * sm_count + 2) // 3) + assert split_count == expected_split_count + else: + assert split_count > 1 + torch.testing.assert_close(output, reference, rtol=3e-2, atol=3e-2) + + +def test_finished_request_publishes_neutral_partials(): + case = _make_case(d=128, p=8, hq=4, requested_splits=8) + finished = torch.tensor([False, True, False], device='cuda') + output, reference, partial_o, partial_ml, split_count = _run( + case, finished) + assert torch.count_nonzero(output[8:16]) == 0 + assert not torch.isnan(output).any() + torch.testing.assert_close(output[:8], reference[:8], rtol=3e-2, atol=3e-2) + torch.testing.assert_close(output[16:], + reference[16:], + rtol=3e-2, + atol=3e-2) + slots = partial_o[:24 * split_count].reshape(24, split_count, 4, 128)[8:16] + ml = partial_ml[:24 * split_count].reshape(24, split_count, 4, 2)[8:16] + assert torch.count_nonzero(slots) == 0 + assert torch.isneginf(ml[..., 0]).all() + assert torch.count_nonzero(ml[..., 1]) == 0 + + direct_cases = ( + _make_case(d=128, p=8, hq=4, requested_splits=1), + _make_case(d=256, p=8, hq=4, requested_splits=1), + _make_case(dtype=torch.float16, + d=256, + p=8, + hq=4, + requested_splits=1), + ) + for direct_case in direct_cases: + (direct_output, direct_reference, direct_partial_o, direct_partial_ml, + direct_split_count) = _run(direct_case, finished) + assert direct_split_count == 1 + assert torch.count_nonzero(direct_output[8:16]) == 0 + assert not torch.isnan(direct_output).any() + torch.testing.assert_close( + direct_output[:8], direct_reference[:8], rtol=3e-2, atol=3e-2) + torch.testing.assert_close( + direct_output[16:], direct_reference[16:], rtol=3e-2, atol=3e-2) + assert torch.isnan(direct_partial_o).all() + assert torch.isnan(direct_partial_ml).all() + + +@pytest.mark.parametrize('case', [ + _make_case(dtype=torch.float16, + d=128, + p=8, + hq=4, + requested_splits=1, + rope_type=1, + rope_dim=128), + _make_case(d=128, + p=8, + hq=4, + rope_type=1, + rope_dim=128, + mrope_mode=1, + mrope_section=(16, 24, 24)), + _make_case(d=128, + p=8, + hq=4, + rope_type=1, + rope_dim=128, + mrope_mode=2, + mrope_section=(16, 24, 24)), +]) +def test_verification_attention_query_transforms(case): + output, reference, _, _, split_count = _run(case) + if case['requested_splits'] == 1: + assert split_count == 1 + torch.testing.assert_close(output, reference, rtol=4e-2, atol=4e-2) diff --git a/tests/turbomind/attention/verification_attention.py b/tests/turbomind/attention/verification_attention.py new file mode 100644 index 0000000000..93db6fa3d4 --- /dev/null +++ b/tests/turbomind/attention/verification_attention.py @@ -0,0 +1,86 @@ +import _turbomind as _tm +import torch + + +def run_verification_attention( + *, + prefix_k, + prefix_v, + prefix_offsets, + packed_qkv, + q_bias, + output, + cache_storage, + block_ptrs, + block_ptr_offsets, + q_offsets, + k_offsets, + finished, + partial_o, + partial_ml, + query_head_count, + kv_head_count, + head_dim, + block_len, + max_query_length, + max_key_length, + window_size=0, + requested_max_split_count=128, + rope_type=0, + rope_dim=0, + rope_base=1_000_000.0, + rope_factor=1.0, + mrope_mode=0, + mrope_section=(0, 0, 0), + mrope_position_ids=None, + mrope_position_delta=None, + mrope_length=None, +): + device = packed_qkv.device + if mrope_position_ids is None: + mrope_position_ids = torch.empty((0, ), + dtype=torch.int32, + device=device) + mrope_position_delta = torch.empty((0, ), + dtype=torch.int32, + device=device) + mrope_length = torch.empty((0, ), dtype=torch.int32, device=device) + return _tm.verification_attention( + prefix_k, + prefix_v, + prefix_offsets, + max( + int(prefix_offsets[i + 1] - prefix_offsets[i]) + for i in range(prefix_offsets.numel() - 1)), + packed_qkv, + q_bias, + output, + cache_storage, + block_ptrs, + block_ptr_offsets, + q_offsets, + k_offsets, + finished, + partial_o, + partial_ml, + query_head_count, + kv_head_count, + head_dim, + block_len, + max_query_length, + max_key_length, + window_size, + requested_max_split_count, + rope_type, + rope_dim, + rope_base, + rope_factor, + mrope_mode, + mrope_section[0], + mrope_section[1], + mrope_section[2], + mrope_position_ids, + mrope_position_delta, + mrope_length, + torch.cuda.current_stream(device).cuda_stream, + ) diff --git a/tests/turbomind/copy/test_copy.py b/tests/turbomind/copy/test_copy.py new file mode 100644 index 0000000000..a652900fab --- /dev/null +++ b/tests/turbomind/copy/test_copy.py @@ -0,0 +1,144 @@ +import itertools +import math + +import pytest +import torch + +tm = pytest.importorskip('_turbomind') +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason='CUDA is required') + + +def _copy(src, dst): + tm.generic_copy_on_stream( + tm.from_dlpack_with_strides(src), + tm.from_dlpack_with_strides(dst), + torch.cuda.current_stream().cuda_stream, + ) + torch.cuda.synchronize() + + +def _check_copy(shape, src_strides, dst_strides, dtype=torch.int32, offset=16): + # Compare the entire backing storage, including padding and guards. + def storage_size(strides): + return 1 + sum((extent - 1) * stride for extent, stride in zip(shape, strides)) + + src_storage = (torch.arange(storage_size(src_strides) + 2 * offset, + device='cuda', dtype=torch.int64) % 97).to(dtype) + dst_storage = torch.full((storage_size(dst_strides) + 2 * offset,), 113, + device='cuda', dtype=dtype) + expected = dst_storage.clone() + src_before = src_storage.clone() + src = src_storage.as_strided(shape, src_strides, offset) + dst = dst_storage.as_strided(shape, dst_strides, offset) + expected.as_strided(shape, dst_strides, offset).copy_(src) + _copy(src, dst) + torch.testing.assert_close(dst_storage, expected, rtol=0, atol=0) + torch.testing.assert_close(src_storage, src_before, rtol=0, atol=0) + + +@pytest.mark.parametrize('dtype', [torch.uint8, torch.int16, torch.int32, torch.int64]) +@pytest.mark.parametrize('layers,requests', [(24, 1), (25, 1), (24, 4), (24, 7)]) +def test_broadcast(dtype, layers, requests): + _check_copy((layers, requests), (0, 1), (requests, 1), dtype) + + +def test_broadcast_to_permuted_destination(): + # Source coalesces its first two axes; destination coalesces its last two. + # Equal resulting ranks do not imply equal resulting coordinate shapes. + _check_copy((2, 3, 4), (0, 0, 1), (1, 8, 2)) + + +@pytest.mark.parametrize('order', list(itertools.permutations(range(3)))) +def test_padded_permutations(order): + shape = (5, 3, 37) + strides = [0] * 3 + pitch = 1 + for axis in order: + strides[axis] = pitch + pitch = pitch * shape[axis] + 1 + _check_copy(shape, (137, 43, 1), tuple(strides), offset=17) + + +@pytest.mark.parametrize('shape,src_strides,dst_strides', [ + ((2, 1, 3, 4), (0, 100, 0, 1), (1, 100, 8, 2)), + ((1, 1, 1), (0, 0, 0), (9, 3, 1)), + ((2, 3, 4, 5, 6), (360, 120, 30, 6, 1), (360, 120, 30, 6, 1)), + ((5, 3, 100), (1, 7, 29), (1, 9, 41)), +]) +def test_strided_and_coalesced_shapes(shape, src_strides, dst_strides): + _check_copy(shape, src_strides, dst_strides) + + +@pytest.mark.parametrize('dtype', [torch.uint8, torch.int16, torch.int32, torch.int64]) +@pytest.mark.parametrize('offset,padding', [(16, 0), (17, 0), (16, 1)]) +def test_transpose_alignment(dtype, offset, padding): + _check_copy((64, 64), (64 + padding, 1), (1, 64 + padding), dtype, offset) + + +@pytest.mark.parametrize('shape,src_strides,dst_strides', [ + ((3, 128, 192), (24576, 192, 1), (24576, 1, 128)), + ((2, 3, 64, 64), (0, 4096, 64, 1), (12288, 4096, 1, 64)), + ((64, 2, 64), (128, 64, 1), (1, 4096, 64)), + ((2, 64, 64), (4161, 64, 1), (4161, 1, 64)), + ((2, 3, 4, 64, 64), (49152, 16384, 4096, 64, 1), + (49152, 16384, 4096, 1, 64)), +]) +def test_batched_transpose(shape, src_strides, dst_strides): + _check_copy(shape, src_strides, dst_strides) + + +@pytest.mark.parametrize('shape', [(0,), (2, 0, 8)]) +def test_empty_copy(shape): + src = torch.empty(shape, device='cuda') + dst = torch.empty_like(src) + _copy(src, dst) + + +def _check_large_copy(shape, dst_strides=None): + count = math.prod(shape) + free, _ = torch.cuda.mem_get_info() + reusable = torch.cuda.memory_reserved() - torch.cuda.memory_allocated() + # Two byte buffers plus space for the comparison and allocator overhead. + if free + reusable < 3 * count + 512 * 1024**2: + pytest.skip('Insufficient free GPU memory for the large-copy regression') + + guard = 64 + src_storage = torch.full((count + 2 * guard,), 251, device='cuda', dtype=torch.uint8) + dst_storage = torch.full_like(src_storage, 251) + src = src_storage[guard:-guard].view(shape) + dst = dst_storage[guard:-guard].view(shape) + if dst_strides is not None: + dst = dst.as_strided(shape, dst_strides) + dst.fill_(253) + + # Generate varying data without a count-sized int64 arange or reference copy. + rows = src.view(-1, shape[-1]) + row_values = torch.arange(rows.shape[0], device='cuda', dtype=torch.int32).remainder_(97).to(torch.uint8) + col_values = torch.arange(rows.shape[1], device='cuda', dtype=torch.int32).mul_(3).remainder_(97).to(torch.uint8) + rows.copy_(row_values[:, None]) + rows.add_(col_values[None, :]) + _copy(src, dst) + + assert torch.equal(dst, src) + for storage in (src_storage, dst_storage): + assert bool((storage[:guard] == 251).all()) + assert bool((storage[-guard:] == 251).all()) + + +@pytest.mark.parametrize('rows', [32767, 32768, 32769]) +def test_coalesced_shape_int32_boundary(rows): + _check_large_copy((rows, 65536)) + + +@pytest.mark.parametrize('rows,columns', [(65535, 64), (32768, 128)]) +def test_coalesced_transpose_grid_y_boundary(rows, columns): + _check_large_copy((64, rows, columns), (1, columns * 64, 64)) + + +@pytest.mark.parametrize('batch', [65535, 65536]) +def test_transpose_grid_z_boundary(batch): + _check_large_copy((batch, 64, 64), (4096, 1, 64)) + + +def test_transpose_large_permuted_batch(): + _check_large_copy((257, 257, 64, 64), (4096, 257 * 4096, 1, 64)) diff --git a/tests/turbomind/draft_carry/__init__.py b/tests/turbomind/draft_carry/__init__.py new file mode 100644 index 0000000000..e412b6756f --- /dev/null +++ b/tests/turbomind/draft_carry/__init__.py @@ -0,0 +1 @@ +"""Native draft-carry kernel tests.""" diff --git a/tests/turbomind/draft_carry/draft_carry.py b/tests/turbomind/draft_carry/draft_carry.py new file mode 100644 index 0000000000..b4fe00a337 --- /dev/null +++ b/tests/turbomind/draft_carry/draft_carry.py @@ -0,0 +1,48 @@ +from __future__ import annotations + +import torch + +_NATIVE_SYMBOL = 'select_draft_carry' + + +def _load_native_bridge(): + try: + import _turbomind as tm + except ImportError: + return None + return tm if hasattr(tm, _NATIVE_SYMBOL) else None + + +def _require_native_bridge(): + tm = _load_native_bridge() + if tm is None: + raise ImportError( + 'TurboMind draft-carry bridge is unavailable; ' + f'required symbol: {_NATIVE_SYMBOL}') + return tm + + +def is_available() -> bool: + return _load_native_bridge() is not None + + +def select_draft_carry( + local_residual: torch.Tensor, + selected_local_rows: torch.Tensor, + candidate_active: torch.Tensor, + carry: torch.Tensor, + first: int, + last: int, +) -> None: + stream_ptr = int( + torch.cuda.current_stream( + local_residual.device).cuda_stream) + _require_native_bridge().select_draft_carry( + local_residual, + selected_local_rows, + candidate_active, + carry, + int(first), + int(last), + stream_ptr, + ) diff --git a/tests/turbomind/draft_carry/reference.py b/tests/turbomind/draft_carry/reference.py new file mode 100644 index 0000000000..2cfdb72306 --- /dev/null +++ b/tests/turbomind/draft_carry/reference.py @@ -0,0 +1,61 @@ +from __future__ import annotations + +from collections.abc import Sequence +from dataclasses import dataclass + +import torch + + +@dataclass(frozen=True) +class OwnedTokenRows: + offset: int + first: int + last: int + + +def compute_token_ownership( + global_rank: int, + tp0: int, + tp1: int, + local_token_nums: Sequence[int], +) -> OwnedTokenRows: + """Compute one rank's packed offset and DP-local model-TP interval.""" + inner_tp = min(tp0, tp1) + dp_index = global_rank // inner_tp + tp_index = global_rank % inner_tp + num = int(local_token_nums[dp_index]) + + slice_size = (num + inner_tp - 1) // inner_tp + first = min(num, tp_index * slice_size) + last = min(num, first + slice_size) + offset = sum(int(value) for value in local_token_nums[:dp_index]) + return OwnedTokenRows(offset=offset, first=first, last=last) + + +def select_draft_carry_reference( + local_residual_bytes: torch.Tensor, + selected_local_rows: Sequence[int], + candidate_active: Sequence[bool], + first: int, + last: int, +) -> torch.Tensor: + """Apply owner-or-zero selection to an opaque two-dimensional byte view.""" + candidate_count = len(selected_local_rows) + row_bytes = local_residual_bytes.shape[1] + carry = torch.zeros( + (candidate_count, row_bytes), + dtype=torch.uint8, + device=local_residual_bytes.device, + ) + + local_token_num = local_residual_bytes.shape[0] + for candidate, (row, active) in enumerate( + zip(selected_local_rows, candidate_active) + ): + if ( + active + and 0 <= row < local_token_num + and first <= row < last + ): + carry[candidate].copy_(local_residual_bytes[row]) + return carry diff --git a/tests/turbomind/draft_carry/test_draft_carry.py b/tests/turbomind/draft_carry/test_draft_carry.py new file mode 100644 index 0000000000..2deee02527 --- /dev/null +++ b/tests/turbomind/draft_carry/test_draft_carry.py @@ -0,0 +1,326 @@ +from __future__ import annotations + +from collections.abc import Sequence + +import pytest +import torch + +from .reference import ( + OwnedTokenRows, + compute_token_ownership, + select_draft_carry_reference, +) +from .draft_carry import ( + is_available, + select_draft_carry, +) + +pytestmark = pytest.mark.skipif( + not torch.cuda.is_available(), + reason='CUDA is required for draft-carry tests', +) + +_HIDDEN_SIZE = 4096 +_GUARD_BYTES = 64 +_GUARD_VALUE = 0xD7 +_OUTPUT_POISON = 0xA5 +_INVALID_ROW = 1_000_000_000 + +_WIDTH_DTYPES = { + 8: (torch.uint8, torch.int8), + 16: (torch.int16, torch.float16), + 32: (torch.int32, torch.float32), +} + +_TOPOLOGIES = [ + pytest.param(1, (4, 3, 0), id='model_tp_1'), + pytest.param(2, (4, 3, 1), id='model_tp_2'), + pytest.param(8, (16, 5, 0), id='model_tp_8'), +] + + +def _require_bridge() -> None: + if not is_available(): + pytest.skip('TurboMind draft-carry bridge is unavailable') + + +def _make_packed_diagnostic( + local_token_nums: Sequence[int], + row_bytes: int, +) -> tuple[torch.Tensor, torch.Tensor]: + total_rows = sum(local_token_nums) + payload_bytes = total_rows * row_bytes + storage = torch.full( + (2 * _GUARD_BYTES + payload_bytes,), + _GUARD_VALUE, + dtype=torch.uint8, + device='cuda', + ) + payload = storage[ + _GUARD_BYTES:_GUARD_BYTES + payload_bytes + ].reshape(total_rows, row_bytes) + if payload_bytes: + values = ( + torch.arange( + payload_bytes, + dtype=torch.int64, + device='cuda', + ) + .mul_(37) + .add_(17) + .remainder_(251) + .add_(1) + .to(torch.uint8) + ) + payload.copy_(values.reshape_as(payload)) + return storage, payload + + +def _guarded_carry( + candidate_count: int, + row_bytes: int, + carry_dtype: torch.dtype, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + payload_bytes = candidate_count * row_bytes + storage = torch.full( + (2 * _GUARD_BYTES + payload_bytes,), + _OUTPUT_POISON, + dtype=torch.uint8, + device='cuda', + ) + payload = storage[ + _GUARD_BYTES:_GUARD_BYTES + payload_bytes + ].reshape(candidate_count, row_bytes) + return storage, payload.view(carry_dtype), payload + + +def _assert_guards( + storage: torch.Tensor, + guard_value: int, +) -> None: + assert torch.all(storage[:_GUARD_BYTES] == guard_value).item() + assert torch.all(storage[-_GUARD_BYTES:] == guard_value).item() + + +def _segment_ownership( + dp_index: int, + model_tp: int, + tp0: int, + tp1: int, + local_token_nums: Sequence[int], +) -> list[OwnedTokenRows]: + ownership = [ + compute_token_ownership( + dp_index * model_tp + tp_index, + tp0, + tp1, + local_token_nums, + ) + for tp_index in range(model_tp) + ] + + num = local_token_nums[dp_index] + offset = sum(local_token_nums[:dp_index]) + cursor = 0 + for owned in ownership: + assert owned.offset == offset + assert owned.first == cursor + assert owned.first <= owned.last <= num + cursor = owned.last + assert cursor == num + return ownership + + +def _candidate_rows( + ownership: Sequence[OwnedTokenRows], + local_token_num: int, +) -> tuple[list[int], list[bool]]: + boundary_rows: list[int] = [] + for owned in ownership: + if owned.first < owned.last: + boundary_rows.extend((owned.first, owned.last - 1)) + if owned.first > 0: + boundary_rows.append(owned.first - 1) + if owned.last < local_token_num: + boundary_rows.append(owned.last) + + selected_rows = list(dict.fromkeys(boundary_rows)) + active = [True] * len(selected_rows) + + inactive_valid = selected_rows[0] if selected_rows else 0 + selected_rows.extend((inactive_valid, _INVALID_ROW, -_INVALID_ROW)) + active.extend((False, False, True)) + return selected_rows, active + + +@pytest.mark.parametrize('element_bits', [8, 16, 32]) +@pytest.mark.parametrize('global_communicator_first', [True, False]) +@pytest.mark.parametrize('model_tp,local_token_nums', _TOPOLOGIES) +def test_model_tp_owner_or_zero_matrix( + model_tp, + local_token_nums, + global_communicator_first, + element_bits, +): + _require_bridge() + global_rank_count = model_tp * len(local_token_nums) + if global_communicator_first: + tp0, tp1 = global_rank_count, model_tp + else: + tp0, tp1 = model_tp, global_rank_count + + source_dtype, carry_dtype = _WIDTH_DTYPES[element_bits] + element_bytes = element_bits // 8 + row_bytes = _HIDDEN_SIZE * element_bytes + diagnostic_storage, diagnostic_bytes = _make_packed_diagnostic( + local_token_nums, + row_bytes, + ) + diagnostic_before = diagnostic_storage.clone() + + for dp_index, local_token_num in enumerate(local_token_nums): + ownership = _segment_ownership( + dp_index, + model_tp, + tp0, + tp1, + local_token_nums, + ) + selected_rows, active = _candidate_rows( + ownership, + local_token_num, + ) + selected_local_rows = torch.tensor( + selected_rows, + dtype=torch.int32, + device='cuda', + ) + candidate_active = torch.tensor( + active, + dtype=torch.bool, + device='cuda', + ) + + offset = ownership[0].offset + local_bytes = diagnostic_bytes[ + offset:offset + local_token_num + ] + local_residual = local_bytes.view(source_dtype) + assert local_residual.shape == ( + local_token_num, + _HIDDEN_SIZE, + ) + + rank_results = [] + for owned in ownership: + carry_storage, carry, carry_bytes = _guarded_carry( + len(selected_rows), + row_bytes, + carry_dtype, + ) + expected = select_draft_carry_reference( + local_bytes, + selected_rows, + active, + owned.first, + owned.last, + ) + + result = select_draft_carry( + local_residual, + selected_local_rows, + candidate_active, + carry, + owned.first, + owned.last, + ) + + assert result is None + assert torch.equal(carry_bytes, expected) + _assert_guards(carry_storage, _OUTPUT_POISON) + rank_results.append(carry_bytes.clone()) + + owner_sum = torch.stack( + rank_results, + dim=0, + ).to(torch.int64).sum(dim=0) + unsharded = select_draft_carry_reference( + local_bytes, + selected_rows, + active, + 0, + local_token_num, + ) + assert torch.equal( + owner_sum, + unsharded.to(torch.int64), + ) + + assert torch.equal(diagnostic_storage, diagnostic_before) + _assert_guards(diagnostic_storage, _GUARD_VALUE) + + +@pytest.mark.parametrize('element_bits', [8, 16, 32]) +def test_identity_selection_supports_exact_payload_alias(element_bits): + _require_bridge() + source_dtype, carry_dtype = _WIDTH_DTYPES[element_bits] + element_bytes = element_bits // 8 + row_bytes = _HIDDEN_SIZE * element_bytes + candidate_count = 6 + payload_bytes = candidate_count * row_bytes + + storage = torch.full( + (2 * _GUARD_BYTES + payload_bytes,), + _GUARD_VALUE, + dtype=torch.uint8, + device='cuda', + ) + payload = storage[ + _GUARD_BYTES:_GUARD_BYTES + payload_bytes + ].reshape(candidate_count, row_bytes) + values = ( + torch.arange( + payload_bytes, + dtype=torch.int64, + device='cuda', + ) + .mul_(29) + .add_(11) + .remainder_(251) + .add_(1) + .to(torch.uint8) + ) + payload.copy_(values.reshape_as(payload)) + + selected_rows = list(range(candidate_count)) + active = [True, True, False, True, True, True] + expected = select_draft_carry_reference( + payload.clone(), + selected_rows, + active, + first=1, + last=5, + ) + selected_local_rows = torch.arange( + candidate_count, + dtype=torch.int32, + device='cuda', + ) + candidate_active = torch.tensor( + active, + dtype=torch.bool, + device='cuda', + ) + + result = select_draft_carry( + payload.view(source_dtype), + selected_local_rows, + candidate_active, + payload.view(carry_dtype), + first=1, + last=5, + ) + + assert result is None + assert torch.equal(payload, expected) + _assert_guards(storage, _GUARD_VALUE) diff --git a/tests/turbomind/linear_attn/benchmark.py b/tests/turbomind/linear_attn/benchmark.py index 435e9f89a0..5d7866c75a 100644 --- a/tests/turbomind/linear_attn/benchmark.py +++ b/tests/turbomind/linear_attn/benchmark.py @@ -4,6 +4,7 @@ import importlib import math import random +import statistics from collections.abc import Callable, Iterable from dataclasses import dataclass, replace from types import ModuleType, SimpleNamespace @@ -27,6 +28,7 @@ ) VALID_BACKENDS = ('reference', 'turbomind', 'fla', 'flashqla') +VALID_GDR_MODES = ('auto', 'recurrent', 'chunked', 'verify') VALID_CP_LEVELS = ('all', 'exact', 'off') VALID_CP_PATTERNS = ('auto', 'warmup', 'fallback', 'alternating') CP_FALLBACK_SEGMENT_LOG_DECAY = -5.0 @@ -41,6 +43,8 @@ class BenchmarkRequest: backend: str state_dtype: str chunk_size: int | None + gdr_mode: str + suppress_state_store: bool cp_level: str cp_pattern: str validate_outputs: bool @@ -251,7 +255,7 @@ def time_task(task: BenchmarkTask, request: BenchmarkRequest, device: torch.devi flusher = L2CacheFlusher(device) if request.l2_flush else None start_event = torch.cuda.Event(enable_timing=True) end_event = torch.cuda.Event(enable_timing=True) - elapsed_ms = 0.0 + latency_samples_ms = [] for _ in range(request.iters): task.prepare() @@ -263,10 +267,11 @@ def time_task(task: BenchmarkTask, request: BenchmarkRequest, device: torch.devi task.run() end_event.record(stream) end_event.synchronize() - elapsed_ms += start_event.elapsed_time(end_event) + latency_samples_ms.append(start_event.elapsed_time(end_event)) row = dict(task.row) - row['latency_ms'] = elapsed_ms / max(request.iters, 1) + row['latency_ms'] = sum(latency_samples_ms) / len(latency_samples_ms) + row['latency_median_ms'] = statistics.median(latency_samples_ms) row['l2_flush_bytes'] = 0 if flusher is None else flusher.bytes if request.print_diffs and validation_row is not None: row.update(validation_row) @@ -894,6 +899,8 @@ def request_from_args(args, *, backend: str, run: RunCase) -> BenchmarkRequest: backend=backend, state_dtype=run.state_dtype, chunk_size=args.chunk_size, + gdr_mode=args.gdr_mode, + suppress_state_store=args.suppress_state_store, cp_level=args.cp_level, cp_pattern=args.cp_pattern, validate_outputs=not args.skip_validate, @@ -910,9 +917,9 @@ def parse_chunk_size(value: str) -> int | None: try: chunk_size = int(value) except ValueError as exc: - raise argparse.ArgumentTypeError('chunk size must be auto, 1, 16, 32, or 64') from exc - if chunk_size not in (1, 16, 32, 64): - raise argparse.ArgumentTypeError('chunk size must be auto, 1, 16, 32, or 64') + raise argparse.ArgumentTypeError('chunk size must be auto, 1, 8, 16, 32, or 64') from exc + if chunk_size not in (1, 8, 16, 32, 64): + raise argparse.ArgumentTypeError('chunk size must be auto, 1, 8, 16, 32, or 64') return chunk_size @@ -954,6 +961,8 @@ def make_parser() -> argparse.ArgumentParser: parser.add_argument('--batch-size', type=int, default=1) parser.add_argument('--backend', default='reference') parser.add_argument('--chunk-size', type=parse_chunk_size, default='auto') + parser.add_argument('--gdr-mode', choices=VALID_GDR_MODES, default='auto') + parser.add_argument('--suppress-state-store', action='store_true') parser.add_argument('--cp-level', choices=VALID_CP_LEVELS, default='all') parser.add_argument('--cp-pattern', choices=VALID_CP_PATTERNS, default='auto') parser.add_argument('--print-diffs', action='store_true') diff --git a/tests/turbomind/linear_attn/test_gated_delta_rule.py b/tests/turbomind/linear_attn/test_gated_delta_rule.py index 02086ba56e..2e129b4f94 100644 --- a/tests/turbomind/linear_attn/test_gated_delta_rule.py +++ b/tests/turbomind/linear_attn/test_gated_delta_rule.py @@ -159,6 +159,91 @@ def _canonical_native_inputs(inputs: InputTensors) -> InputTensors: return native +@cuda_required +@pytest.mark.skipif(_device_capability() != (9, 0), reason='SM90 is required') +@pytest.mark.parametrize('positions,requested_capacity,expected_capacity', [ + (4, 8, 8), + (8, 8, 8), + (12, None, 16), + (16, 16, 16), +]) +@pytest.mark.parametrize('state_dtype', ['f32', 'bf16']) +def test_sm90_verify_gdr(positions, requested_capacity, expected_capacity, state_dtype): + case = InputCase( + layout=Fixed(batch_size=3, seq_len=positions), + heads=Heads(hq=4, hv=8), + input_dtype=torch.bfloat16, + has_h0=True, + seed=81000 + positions, + ) + run = RunCase(input=case, state_dtype=state_dtype, chunk_size=expected_capacity) + inputs = make_input_tensors(case, device='cuda') + native = _canonical_native_inputs(inputs) + state = make_state_buffer(inputs.h0, run, torch.device('cuda')) + entry_state = state.storage.clone() + finished = torch.ones(3, device='cuda', dtype=torch.bool) + native_state_dtype = _native_state_dtype(state_dtype) + bridge = turbomind_gated_delta_rule.NativeBridge(turbomind_gated_delta_rule._require_native_bridge()) + verify_plan = bridge.plan( + native, + q_offsets=None, + state_dtype=native_state_dtype, + mode='verify', + chunk_size=requested_capacity, + cp_level='off', + num_head_groups=1, + heads_per_block=8, + ) + assert verify_plan['kernel']['mode'] == 'verify' + assert verify_plan['kernel']['chunk_size'] == expected_capacity + assert verify_plan['problem']['chunk_size'] == expected_capacity + + expected, _ = chunk_gated_delta_rule_fwd( + inputs.q, + inputs.k, + inputs.v, + inputs.g, + inputs.beta, + initial_state=entry_state.float(), + chunk_size=64, + ) + dedicated = turbomind_gated_delta_rule.chunk_gated_delta_rule_fwd( + native.q, + native.k, + native.v, + native.g, + native.beta, + state_ptrs=state.ptrs, + finished=finished, + state_dtype=native_state_dtype, + plan=verify_plan, + cp_level='off', + num_head_groups=1, + heads_per_block=8, + ) + torch.testing.assert_close(dedicated, expected, rtol=8e-2, atol=8e-2) + assert torch.equal(state.storage, entry_state) + + state.reset(inputs.h0) + legacy = turbomind_gated_delta_rule.chunk_gated_delta_rule_fwd( + native.q, + native.k, + native.v, + native.g, + native.beta, + state_ptrs=state.ptrs, + finished=finished, + state_dtype=native_state_dtype, + mode='chunked', + chunk_size=64, + cp_level='off', + num_head_groups=1, + heads_per_block=8, + ) + torch.testing.assert_close(legacy, expected, rtol=8e-2, atol=8e-2) + assert torch.equal(state.storage, entry_state) + + def _supported_arch_cases(): if not torch.cuda.is_available(): return [] @@ -187,6 +272,361 @@ def _supported_arch_cases(): return cases +def _transaction_bridge(): + symbols = ( + *turbomind_gated_delta_rule.REQUIRED_NATIVE_BRIDGE_SYMBOLS, + 'gdn_build_state_store_mask', + 'gdn_capture_transitions', + 'gdn_commit_conv_state', + 'gdn_commit_recurrent_state', + ) + return turbomind_gated_delta_rule.NativeBridge( + turbomind_gated_delta_rule._require_native_bridge(symbols)) + + +@cuda_required +def test_speculative_state_mask(): + bridge = _transaction_bridge() + finished = torch.tensor( + [False, True, False, True, False], device='cuda', dtype=torch.bool) + speculative = torch.tensor( + [False, False, True, True, False], device='cuda', dtype=torch.bool) + actual = torch.empty_like(finished) + bridge.build_state_store_mask(actual, finished, speculative) + torch.testing.assert_close(actual, finished | speculative) + + +@cuda_required +@pytest.mark.parametrize('position_count', [2, 4]) +@pytest.mark.parametrize('tp', [1, 2, 4, 8, 16]) +def test_speculative_transition_capture(position_count, tp): + bridge = _transaction_bridge() + dtype = torch.bfloat16 + hq = 16 // tp + hv = 48 // tp + conv_dim = (16 + 48) * 128 // tp + value_dim = 48 * 128 // tp + all_proj_width = conv_dim + value_dim + 2 * hv + + lengths = [1, position_count + 1, 2, position_count] + q_offsets = torch.tensor( + [0, *torch.tensor(lengths).cumsum(0).tolist()], + device='cuda', dtype=torch.int32) + requests = torch.tensor([1, 3], device='cuda', dtype=torch.int32) + token_count = sum(lengths) + poison = 31744.0 + + all_proj = torch.full( + (token_count, all_proj_width), poison, device='cuda', dtype=dtype) + raw = all_proj[:, :conv_dim] + raw.copy_(torch.arange( + token_count * conv_dim, device='cuda', dtype=torch.float32) + .reshape(token_count, conv_dim).remainder(97).to(dtype)) + + key_storage = torch.full( + (1, token_count, hq, 132), poison, device='cuda', dtype=dtype) + value_storage = torch.full( + (1, token_count, hv, 132), poison, device='cuda', dtype=dtype) + key = key_storage[..., :128] + value = value_storage[..., :128] + key.copy_(torch.randn_like(key)) + value.copy_(torch.randn_like(value)) + + decay_storage = torch.full( + (1, token_count, hv + 3), poison, device='cuda', dtype=torch.float32) + beta_storage = torch.full_like(decay_storage, poison) + decay = decay_storage[..., :hv] + beta = beta_storage[..., :hv] + decay.copy_(torch.randn_like(decay) * .01) + beta.copy_(torch.sigmoid(torch.randn_like(beta))) + + layers = 2 + spec_count = requests.numel() + journal_raw = torch.full( + (layers, spec_count, position_count, conv_dim), poison, + device='cuda', dtype=dtype) + journal_key = torch.full( + (layers, spec_count, position_count, hq, 128), poison, + device='cuda', dtype=dtype) + journal_value = torch.full( + (layers, spec_count, position_count, hv, 128), poison, + device='cuda', dtype=dtype) + journal_decay = torch.full( + (layers, spec_count, position_count, hv), poison, + device='cuda', dtype=torch.float32) + journal_beta = torch.full_like(journal_decay, poison) + journal = ( + journal_raw, journal_key, journal_value, journal_decay, journal_beta) + + for layer in range(layers): + bridge.capture_transitions( + raw, key, value, decay, beta, q_offsets, requests, + gdn_layer=layer, + verify_positions=position_count, + journal=journal) + + starts = q_offsets[requests].cpu().tolist() + expected_rows = torch.cat([ + torch.arange(start, start + position_count, device='cuda') + for start in starts + ]) + expected_raw = raw[expected_rows].reshape(spec_count, position_count, conv_dim) + expected_key = key[0, expected_rows].reshape(spec_count, position_count, hq, 128) + expected_value = value[0, expected_rows].reshape(spec_count, position_count, hv, 128) + expected_decay = decay[0, expected_rows].reshape(spec_count, position_count, hv) + expected_beta = beta[0, expected_rows].reshape(spec_count, position_count, hv) + for layer in range(layers): + torch.testing.assert_close(journal_raw[layer], expected_raw) + torch.testing.assert_close(journal_key[layer], expected_key) + torch.testing.assert_close(journal_value[layer], expected_value) + torch.testing.assert_close(journal_decay[layer], expected_decay) + torch.testing.assert_close(journal_beta[layer], expected_beta) + + +@cuda_required +def test_speculative_state_commit_conv_ring_and_guards(): + bridge = _transaction_bridge() + layers, batch, spec_count, position_count = 2, 4, 2, 4 + conv_dim, d_conv = 19, 4 + request_indices = torch.tensor([1, 3], device='cuda', dtype=torch.int32) + entry = torch.tensor([0, 3, 0, 6], device='cuda', dtype=torch.int32) + accept = torch.tensor([0, 0, 0, 3], device='cuda', dtype=torch.int32) + offsets = torch.tensor( + [layer * d_conv * conv_dim for layer in range(layers)], + device='cuda', dtype=torch.int32) + raw = torch.arange( + layers * spec_count * position_count * conv_dim, + device='cuda', dtype=torch.float32).reshape( + layers, spec_count, position_count, conv_dim).to(torch.bfloat16) + + states = [torch.full( + (layers, d_conv, conv_dim), -7, device='cuda', dtype=torch.bfloat16) + for _ in range(batch)] + before = [state.clone() for state in states] + pointers = torch.tensor( + [state.data_ptr() for state in states], device='cuda', dtype=torch.int64) + bridge.commit_conv_state( + raw, pointers, request_indices, entry, accept, offsets, + conv_dim=conv_dim, d_conv=d_conv) + + torch.testing.assert_close(states[0], before[0]) + torch.testing.assert_close(states[1], before[1]) + torch.testing.assert_close(states[2], before[2]) + expected = before[3] + for layer in range(layers): + for position in range(accept[3].item()): + expected[layer, (entry[3].item() - 1 + position) % d_conv] = raw[layer, 1, position] + torch.testing.assert_close(states[3], expected) + + +def _commit_reference(initial, key, value, decay, beta, request_indices, accept_len, + layers_per_block, heads_per_block): + result = initial.float().clone() + _, _, _, hq, _ = key.shape + hv = value.shape[3] + for layer in range(key.shape[0]): + layer_group, layer_in_group = divmod(layer, layers_per_block) + for compact, request in enumerate(request_indices.tolist()): + for value_head in range(hv): + head_group, local_head = divmod(value_head, heads_per_block) + state = result[layer_group, request, head_group, + layer_in_group, local_head] + key_head = value_head // (hv // hq) + for position in range(int(accept_len[request])): + state.mul_(decay[layer, compact, position, value_head].exp()) + key_row = key[layer, compact, position, key_head].float() + value_row = value[layer, compact, position, value_head].float() + prediction = key_row @ state + delta = (value_row - prediction) * beta[ + layer, compact, position, value_head] + state.add_(key_row[:, None] * delta[None, :]) + return result + + +def _state_pointers(storage): + layer_groups, batch, head_groups = storage.shape[:3] + pointers = torch.empty( + layer_groups, batch, head_groups, + device=storage.device, dtype=torch.int64) + values = [ + storage[group, request, head_group].data_ptr() + for group in range(layer_groups) + for request in range(batch) + for head_group in range(head_groups) + ] + pointers.copy_(torch.tensor(values, dtype=torch.int64).view_as(pointers)) + return pointers + + +def _storage_spacing(reference, state_dtype): + stored = reference.to(state_dtype) + next_up = torch.nextafter(stored, torch.full_like(stored, float('inf'))) + next_down = torch.nextafter(stored, torch.full_like(stored, -float('inf'))) + return torch.maximum( + (next_up.float() - stored.float()).abs(), + (stored.float() - next_down.float()).abs()) + + +@cuda_required +@pytest.mark.parametrize('position_count', [2, 4]) +@pytest.mark.parametrize('state_dtype', [torch.bfloat16, torch.float32]) +def test_speculative_state_commit_differential(position_count, state_dtype): + bridge = _transaction_bridge() + torch.manual_seed(91023 + position_count) + layers, batch, spec_count = 1, 2, 2 + hq, hv = 2, 4 + layers_per_block, heads_per_block = 1, 2 + layer_groups = (layers + layers_per_block - 1) // layers_per_block + head_groups = hv // heads_per_block + request_indices = torch.arange(batch, device='cuda', dtype=torch.int32) + key = torch.randn( + layers, spec_count, position_count, hq, 128, + device='cuda', dtype=torch.bfloat16) * .04 + value = torch.randn( + layers, spec_count, position_count, hv, 128, + device='cuda', dtype=torch.bfloat16) * .04 + decay = -torch.rand( + layers, spec_count, position_count, hv, + device='cuda', dtype=torch.float32) * .03 + beta = torch.rand_like(decay) * .2 + + for accepted in range(position_count + 1): + accept_len = torch.tensor( + [accepted, position_count - accepted], + device='cuda', dtype=torch.int32) + initial = torch.randn( + layer_groups, batch, head_groups, layers_per_block, + heads_per_block, 128, 128, + device='cuda', dtype=state_dtype) * .01 + + transactional = initial.clone() + transaction_ptrs = _state_pointers(transactional) + bridge.commit_recurrent_state( + key, value, decay, beta, transaction_ptrs, request_indices, accept_len, + state_dtype='bf16' if state_dtype == torch.bfloat16 else 'f32', + layers_per_block=layers_per_block, + heads_per_block=heads_per_block) + + recurrent = initial.clone() + recurrent_ptrs = _state_pointers(recurrent) + for position in range(position_count): + finished = accept_len <= position + turbomind_gated_delta_rule.chunk_gated_delta_rule_fwd( + key[0, :, position].unsqueeze(1), + key[0, :, position].unsqueeze(1), + value[0, :, position].unsqueeze(1), + decay[0, :, position].unsqueeze(1), + beta[0, :, position].unsqueeze(1), + state_ptrs=recurrent_ptrs[0], + finished=finished, + state_dtype='bf16' if state_dtype == torch.bfloat16 else 'f32', + mode='recurrent', + cp_level='off', + num_head_groups=head_groups, + heads_per_block=heads_per_block, + layer_groups=1, + layers_per_block=1) + + chunked = initial.clone() + chunked_ptrs = _state_pointers(chunked) + for request, accepted_positions in enumerate(accept_len.tolist()): + if accepted_positions == 0: + continue + request_slice = slice(request, request + 1) + position_slice = slice(0, accepted_positions) + turbomind_gated_delta_rule.chunk_gated_delta_rule_fwd( + key[0, request_slice, position_slice], + key[0, request_slice, position_slice], + value[0, request_slice, position_slice], + decay[0, request_slice, position_slice], + beta[0, request_slice, position_slice], + state_ptrs=chunked_ptrs[0, request_slice], + finished=torch.zeros(1, device='cuda', dtype=torch.bool), + state_dtype='bf16' if state_dtype == torch.bfloat16 else 'f32', + mode='chunked', + cp_level='off', + num_head_groups=head_groups, + heads_per_block=heads_per_block, + layer_groups=1, + layers_per_block=1) + + reference = _commit_reference( + initial, key, value, decay, beta, request_indices.cpu(), + accept_len.cpu(), layers_per_block, heads_per_block) + rtol = 8e-2 + atol = 8e-2 + recurrent_f32 = recurrent.float() + chunked_f32 = chunked.float() + transactional_f32 = transactional.float() + torch.testing.assert_close(recurrent_f32, reference, rtol=rtol, atol=atol) + torch.testing.assert_close(chunked_f32, reference, rtol=rtol, atol=atol) + torch.testing.assert_close(transactional_f32, reference, rtol=rtol, atol=atol) + + recurrent_error = (recurrent_f32 - reference).abs() + chunked_error = (chunked_f32 - reference).abs() + transaction_error = (transactional_f32 - reference).abs() + backend_envelope = torch.maximum(recurrent_error, chunked_error) + roundoff = _storage_spacing(reference, state_dtype) + assert torch.all(transaction_error <= backend_envelope + roundoff) + + for request, accepted_positions in enumerate(accept_len.tolist()): + if accepted_positions == 0: + assert torch.equal(transactional[:, request], initial[:, request]) + + +@cuda_required +def test_speculative_state_commit_rounds_only_at_final_store(): + bridge = _transaction_bridge() + position_count = 4 + key = torch.zeros( + 1, 1, position_count, 1, 128, + device='cuda', dtype=torch.bfloat16) + key[..., 0] = 1 + value = torch.full( + (1, 1, position_count, 1, 128), 2, + device='cuda', dtype=torch.bfloat16) + decay = torch.zeros( + 1, 1, position_count, 1, + device='cuda', dtype=torch.float32) + beta = torch.full_like(decay, 1 / 256) + request_indices = torch.zeros(1, device='cuda', dtype=torch.int32) + accept_len = torch.full( + (1,), position_count, + device='cuda', dtype=torch.int32) + initial = torch.zeros( + 1, 1, 1, 1, 1, 128, 128, + device='cuda', dtype=torch.bfloat16) + initial[..., 0, :] = 1 + + transactional = initial.clone() + bridge.commit_recurrent_state( + key, value, decay, beta, _state_pointers(transactional), + request_indices, accept_len, + state_dtype='bf16', layers_per_block=1, heads_per_block=1) + + recurrent = initial.clone() + recurrent_ptrs = _state_pointers(recurrent) + for position in range(position_count): + turbomind_gated_delta_rule.chunk_gated_delta_rule_fwd( + key[0, :, position].unsqueeze(1), + key[0, :, position].unsqueeze(1), + value[0, :, position].unsqueeze(1), + decay[0, :, position].unsqueeze(1), + beta[0, :, position].unsqueeze(1), + state_ptrs=recurrent_ptrs[0], + finished=torch.zeros(1, device='cuda', dtype=torch.bool), + state_dtype='bf16', mode='recurrent', cp_level='off', + num_head_groups=1, heads_per_block=1, + layer_groups=1, layers_per_block=1) + + reference = _commit_reference( + initial, key, value, decay, beta, request_indices.cpu(), + accept_len.cpu(), layers_per_block=1, heads_per_block=1) + final_store_reference = reference.to(torch.bfloat16) + assert not torch.equal(recurrent, final_store_reference) + assert torch.equal(transactional, final_store_reference) + + @cuda_required @pytest.mark.parametrize( 'input_dtype,state_dtype,chunk_size,mode', diff --git a/tests/turbomind/linear_attn/turbomind_gated_delta_rule.py b/tests/turbomind/linear_attn/turbomind_gated_delta_rule.py index b722c71477..9d198aaeae 100644 --- a/tests/turbomind/linear_attn/turbomind_gated_delta_rule.py +++ b/tests/turbomind/linear_attn/turbomind_gated_delta_rule.py @@ -102,6 +102,93 @@ def run(self, tensors, *, plan, **kwargs): finished=self.tensor(kwargs.get('finished')), ) + def build_state_store_mask(self, out, finished, speculative): + self.tm.gdn_build_state_store_mask( + self.tensor(out), + self.tensor(finished), + self.tensor(speculative), + stream_ptr=_current_stream_ptr(out), + ) + + def capture_transitions( + self, + raw_projection, + normalized_key, + value, + log_decay, + beta, + q_offsets, + request_indices, + *, + gdn_layer, + verify_positions, + journal, + ): + self.tm.gdn_capture_transitions( + self.tensor(raw_projection), + self.tensor(normalized_key), + self.tensor(value), + self.tensor(log_decay), + self.tensor(beta), + self.tensor(q_offsets), + self.tensor(request_indices), + gdn_layer, + verify_positions, + *(self.tensor(tensor) for tensor in journal), + stream_ptr=_current_stream_ptr(raw_projection), + ) + + def commit_conv_state( + self, + raw_conv, + conv_state_ptrs, + request_indices, + entry_sequence_length, + accept_len, + conv_state_offsets, + *, + conv_dim, + d_conv, + ): + self.tm.gdn_commit_conv_state( + self.tensor(raw_conv), + self.tensor(conv_state_ptrs), + self.tensor(request_indices), + self.tensor(entry_sequence_length), + self.tensor(accept_len), + self.tensor(conv_state_offsets), + conv_dim, + d_conv, + stream_ptr=_current_stream_ptr(raw_conv), + ) + + def commit_recurrent_state( + self, + key, + value, + log_decay, + beta, + recurrent_state_ptrs, + request_indices, + accept_len, + *, + state_dtype, + layers_per_block, + heads_per_block, + ): + self.tm.gdn_commit_recurrent_state( + self.tensor(key), + self.tensor(value), + self.tensor(log_decay), + self.tensor(beta), + self.tensor(recurrent_state_ptrs), + self.tensor(request_indices), + self.tensor(accept_len), + state_dtype, + layers_per_block, + heads_per_block, + stream_ptr=_current_stream_ptr(key), + ) def validate_benchmark_case(run: RunCase, request: BenchmarkRequest) -> None: if run.input.input_dtype != torch.bfloat16: @@ -186,10 +273,16 @@ def _state_kwargs( state, state_dtype, chunk_size: int | None = None, + *, + recurrent: bool | None = None, + suppress_state_store: bool = False, ): del inputs, bridge, state_dtype - if _is_recurrent_run(run, chunk_size): - finished = torch.zeros(run.input.real_batch_size, device=device, dtype=torch.bool) + if recurrent is None: + recurrent = _is_recurrent_run(run, chunk_size) + if recurrent: + finished = torch.full( + (run.input.real_batch_size,), suppress_state_store, device=device, dtype=torch.bool) return native_inputs, { 'state_ptrs': state.ptrs[:, None], 'finished': finished, @@ -197,7 +290,7 @@ def _state_kwargs( }, True q_offsets = _q_offsets_for_tensors(native_inputs) sequence_num = _sequence_num(run.input, q_offsets) - finished = torch.zeros(sequence_num, device=device, dtype=torch.bool) + finished = torch.full((sequence_num,), suppress_state_store, device=device, dtype=torch.bool) kwargs = {'state_ptrs': state.ptrs, 'finished': finished} if q_offsets is not None: kwargs['q_offsets'] = q_offsets @@ -247,17 +340,18 @@ def chunk_gated_delta_rule_fwd( ) -> torch.Tensor: bridge = NativeBridge(_require_native_bridge()) if plan is not None: + effective_mode = plan['kernel']['mode'] effective_chunk_size = int(plan['problem']['chunk_size']) - recurrent = plan['kernel']['mode'] == 'recurrent' elif mode is not None: - recurrent = mode == 'recurrent' - effective_chunk_size = 1 if recurrent and chunk_size is None else chunk_size + effective_mode = mode + effective_chunk_size = 1 if mode == 'recurrent' and chunk_size is None else chunk_size else: effective_chunk_size = chunk_size recurrent = effective_chunk_size == 1 or ( effective_chunk_size is None and q_offsets is None and q.dtype == torch.bfloat16 and q.shape[1] == 1 ) - if q_offsets is None and not recurrent: + effective_mode = 'recurrent' if recurrent else 'chunked' + if q_offsets is None and effective_mode == 'chunked': q_offsets = torch.arange(0, (q.shape[0] + 1) * q.shape[1], q.shape[1], device=q.device, dtype=torch.int32) tensors = InputTensors(q=q, k=k, v=v, g=g, beta=beta, h0=None, offsets=q_offsets) @@ -266,15 +360,17 @@ def chunk_gated_delta_rule_fwd( tensors, q_offsets=q_offsets, state_dtype=state_dtype, - mode='recurrent' if recurrent else 'chunked', - chunk_size=chunk_size, + mode=effective_mode, + chunk_size=effective_chunk_size, cp_level=cp_level, num_head_groups=num_head_groups, heads_per_block=heads_per_block or v.shape[2], ) - execution_state_ptrs = state_ptrs[:, None] if recurrent and state_ptrs.ndim == 1 else state_ptrs - if recurrent and state_tma_descs is None and plan['kernel']['arch'] != 'pre_sm90': + needs_2d_state_ptrs = effective_mode in ('recurrent', 'verify') + execution_state_ptrs = state_ptrs[:, None] if needs_2d_state_ptrs and state_ptrs.ndim == 1 else state_ptrs + uses_state_tma = needs_2d_state_ptrs and plan['kernel']['arch'] != 'pre_sm90' + if uses_state_tma and state_tma_descs is None: prepared_descs = torch.empty( (layer_groups, execution_state_ptrs.shape[-2], plan['problem']['num_head_groups'], 128), device=q.device, @@ -317,8 +413,12 @@ def _turbomind_task(run: RunCase, inputs: InputTensors, request: BenchmarkReques native_inputs = make_packed_qkv_views(_native_aligned_tensors(inputs)) state_arg = _state_dtype_arg(request.state_dtype) chunk_size = request.chunk_size + mode = request.gdr_mode + if mode == 'auto': + mode = 'recurrent' if _is_recurrent_run(run, chunk_size) else 'chunked' + recurrent = mode == 'recurrent' state = make_state_buffer(inputs.h0, run, device) - run_inputs, kwargs, recurrent = _state_kwargs( + run_inputs, kwargs, _ = _state_kwargs( run, inputs, native_inputs, @@ -327,13 +427,15 @@ def _turbomind_task(run: RunCase, inputs: InputTensors, request: BenchmarkReques state, request.state_dtype, chunk_size, + recurrent=recurrent, + suppress_state_store=request.suppress_state_store, ) out = torch.empty_like(inputs.v) plan = bridge.plan( run_inputs, q_offsets=kwargs.get('q_offsets'), state_dtype=state_arg, - mode='recurrent' if recurrent else 'chunked', + mode=mode, chunk_size=chunk_size, cp_level=request.cp_level, num_head_groups=1, @@ -361,7 +463,10 @@ def _turbomind_task(run: RunCase, inputs: InputTensors, request: BenchmarkReques ) run_inputs.g.copy_(controlled_inputs.g) inputs = controlled_inputs - if recurrent and plan['kernel']['arch'] != 'pre_sm90': + uses_state_tma = mode in ('recurrent', 'verify') and plan['kernel']['arch'] != 'pre_sm90' + if uses_state_tma and kwargs['state_ptrs'].ndim == 1: + kwargs['state_ptrs'] = kwargs['state_ptrs'][:, None] + if uses_state_tma: prepared_descs = torch.empty( (1, run.input.real_batch_size, 1, 128), device=device, @@ -378,9 +483,12 @@ def _turbomind_task(run: RunCase, inputs: InputTensors, request: BenchmarkReques state.tma_descs = prepared_descs kwargs['state_tma_descs'] = state.tma_descs planned_chunk_size = int(plan['problem']['chunk_size']) - if recurrent: + if mode == 'recurrent': if planned_chunk_size != 1: raise RuntimeError(f'native recurrent plan selected chunk_size={planned_chunk_size}, expected 1') + elif mode == 'verify': + if planned_chunk_size not in (8, 16): + raise RuntimeError(f'native verify plan selected chunk_size={planned_chunk_size}, expected 8 or 16') elif planned_chunk_size <= 1: raise RuntimeError(f'native chunked plan selected chunk_size={planned_chunk_size}, expected > 1') workspace = ( @@ -397,6 +505,7 @@ def _turbomind_task(run: RunCase, inputs: InputTensors, request: BenchmarkReques cp_level=request.cp_level, cp_pattern=request.cp_pattern, cp_enabled=cp_enabled, + gdr_mode=mode, ) def prepare(): state.reset(inputs.h0) @@ -412,6 +521,7 @@ def execute(): q_offsets=kwargs.get('q_offsets'), finished=kwargs['finished'], state_dtype=state_arg, + mode=mode, chunk_size=planned_chunk_size, cp_level=request.cp_level, out=out, @@ -424,7 +534,8 @@ def validate(): prepare() actual_o = execute() torch.cuda.synchronize(device) - expected_o, expected_state = reference_chunk_gated_delta_rule_fwd( + reference_chunk_size = 64 if mode == 'verify' else planned_chunk_size + expected_o, transitioned_state = reference_chunk_gated_delta_rule_fwd( inputs.q, inputs.k, inputs.v, @@ -432,8 +543,10 @@ def validate(): inputs.beta, initial_state=inputs.h0, cu_seqlens=inputs.offsets, - chunk_size=planned_chunk_size, + chunk_size=reference_chunk_size, ) + entry_state = torch.zeros_like(transitioned_state) if inputs.h0 is None else inputs.h0 + expected_state = entry_state if mode == 'verify' or request.suppress_state_store else transitioned_state if request.validate_outputs: torch.testing.assert_close(actual_o, expected_o, rtol=8e-2, atol=8e-2) torch.testing.assert_close(state.storage.float(), expected_state, rtol=8e-2, atol=8e-2) diff --git a/tests/turbomind/speculative_sampling/__init__.py b/tests/turbomind/speculative_sampling/__init__.py new file mode 100644 index 0000000000..035543b17d --- /dev/null +++ b/tests/turbomind/speculative_sampling/__init__.py @@ -0,0 +1 @@ +"""TurboMind speculative-sampling kernel tests.""" diff --git a/tests/turbomind/speculative_sampling/reference.py b/tests/turbomind/speculative_sampling/reference.py new file mode 100644 index 0000000000..fb7bae1208 --- /dev/null +++ b/tests/turbomind/speculative_sampling/reference.py @@ -0,0 +1,52 @@ +from collections.abc import Sequence + +import torch + + +def greedy_reference( + probabilities: torch.Tensor, token_ids: torch.Tensor, + kept_count: torch.Tensor, + draft_token_ids: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """Reference the deterministic first-maximum greedy decision.""" + selected = [] + accepted = [] + probabilities_cpu = probabilities.cpu() + token_ids_cpu = token_ids.cpu() + kept_count_cpu = kept_count.cpu() + draft_token_ids_cpu = draft_token_ids.cpu() + + for row in range(probabilities.shape[0]): + kept = int(kept_count_cpu[row]) + target_index = int(torch.argmax(probabilities_cpu[row, :kept])) + target_token = int(token_ids_cpu[row, target_index]) + draft_token = int(draft_token_ids_cpu[row]) + selected.append(target_token) + accepted.append(target_token == draft_token) + + return (torch.tensor(selected, dtype=torch.int32), + torch.tensor(accepted, dtype=torch.bool)) + + +def recovery_token_reference( + probabilities: torch.Tensor, token_ids: torch.Tensor, + kept_count: torch.Tensor, + draft_token_ids: torch.Tensor) -> Sequence[set[int]]: + """Return the valid positive-mass non-draft corrections for each row.""" + valid_tokens = [] + probabilities_cpu = probabilities.cpu() + token_ids_cpu = token_ids.cpu() + kept_count_cpu = kept_count.cpu() + draft_token_ids_cpu = draft_token_ids.cpu() + + for row in range(probabilities.shape[0]): + kept = int(kept_count_cpu[row]) + draft_token = int(draft_token_ids_cpu[row]) + row_tokens = { + int(token_ids_cpu[row, index]) + for index in range(kept) + if int(token_ids_cpu[row, index]) != draft_token + and float(probabilities_cpu[row, index]) > 0.0 + } + valid_tokens.append(row_tokens) + + return valid_tokens diff --git a/tests/turbomind/speculative_sampling/speculative_sampling.py b/tests/turbomind/speculative_sampling/speculative_sampling.py new file mode 100644 index 0000000000..78e9a5a7e5 --- /dev/null +++ b/tests/turbomind/speculative_sampling/speculative_sampling.py @@ -0,0 +1,108 @@ +import _turbomind as _tm +import torch + +STATE_CAPACITY_PER_LOGICAL_STATE = 4096 + + +def allocate_random_states(random_state_count: int, + device: torch.device | str) -> torch.Tensor: + return torch.full( + (random_state_count * STATE_CAPACITY_PER_LOGICAL_STATE, ), + 0xA5, + dtype=torch.uint8, + device=device) + + +def initialize_random_states(random_states: torch.Tensor, + random_seeds: torch.Tensor, + initialize: torch.Tensor) -> None: + stream_ptr = torch.cuda.current_stream(random_states.device).cuda_stream + _tm.initialize_speculative_sampling_states( + random_states, + random_seeds.numel(), + random_seeds, + initialize, + stream_ptr, + ) + + +def verify_target_block( + probabilities: torch.Tensor, + probability_token_ids: torch.Tensor, + kept_count: torch.Tensor, + verification_draft_ids: torch.Tensor, + greedy: torch.Tensor, + logits_active: torch.Tensor, + random_states: torch.Tensor, + random_state_indices: torch.Tensor, + request_token_ids_ptrs: torch.Tensor, + entry_sequence_length: torch.Tensor, + request_to_generation_offsets: torch.Tensor, + speculative_row: torch.Tensor, + selected_span_ids: torch.Tensor, + accept_len: torch.Tensor, + accepted_draft_count: torch.Tensor | None, + position_count: int, +) -> None: + stream_ptr = torch.cuda.current_stream(probabilities.device).cuda_stream + _tm.verify_target_block( + probabilities, + probability_token_ids, + kept_count, + verification_draft_ids, + greedy, + logits_active, + random_states, + random_state_indices, + request_token_ids_ptrs, + entry_sequence_length, + request_to_generation_offsets, + speculative_row, + selected_span_ids, + accept_len, + accepted_draft_count, + position_count, + stream_ptr, + ) + + +def sample_processed_probabilities( + probabilities: torch.Tensor, + indices: torch.Tensor, + kept: torch.Tensor, + curand_states: torch.Tensor, + curand_state_indices: torch.Tensor, + sample_mask: torch.Tensor | None, + selected_tokens: torch.Tensor, + sampled_logprobs: torch.Tensor | None, + sampled_indexes: torch.Tensor | None, + sampled_nums: torch.Tensor | None, +) -> None: + stream_ptr = torch.cuda.current_stream(probabilities.device).cuda_stream + _tm.sample_processed_probabilities( + probabilities, + indices, + kept, + curand_states, + curand_state_indices, + sample_mask, + selected_tokens, + sampled_logprobs, + sampled_indexes, + sampled_nums, + stream_ptr, + ) + + +def append_one_token_and_advance_sequence( + token_ids_ptrs: torch.Tensor, + selected_tokens: torch.Tensor, + sequence_length: torch.Tensor, +) -> None: + stream_ptr = torch.cuda.current_stream(token_ids_ptrs.device).cuda_stream + _tm.append_one_token_and_advance_sequence( + token_ids_ptrs, + selected_tokens, + sequence_length, + stream_ptr, + ) diff --git a/tests/turbomind/speculative_sampling/test_speculative_sampling.py b/tests/turbomind/speculative_sampling/test_speculative_sampling.py new file mode 100644 index 0000000000..26cac21da9 --- /dev/null +++ b/tests/turbomind/speculative_sampling/test_speculative_sampling.py @@ -0,0 +1,875 @@ +from dataclasses import dataclass + +import pytest +import torch + +from .reference import greedy_reference, recovery_token_reference +from .speculative_sampling import ( + allocate_random_states, + append_one_token_and_advance_sequence, + initialize_random_states, + sample_processed_probabilities, + verify_target_block, +) + +_MAX_LOGPROB = 1024 +_OUTPUT_GUARD = 16 +_SELECTED_GUARD_VALUE = 0x5A5A5A5A + + +@dataclass +class GuardedOutput: + storage: torch.Tensor + value: torch.Tensor + guard_value: int | float | bool + + def assert_guards(self) -> None: + leading = self.storage[:_OUTPUT_GUARD] + trailing = self.storage[-_OUTPUT_GUARD:] + expected_leading = torch.full_like(leading, self.guard_value) + expected_trailing = torch.full_like(trailing, self.guard_value) + assert torch.equal(leading, expected_leading) + assert torch.equal(trailing, expected_trailing) + + +def _guarded_output(batch_size: int, dtype: torch.dtype, + device: torch.device | str, + interior_value: int | float | bool) -> GuardedOutput: + guard_value = True if dtype == torch.bool else _SELECTED_GUARD_VALUE + storage = torch.full((batch_size + 2 * _OUTPUT_GUARD, ), + guard_value, + dtype=dtype, + device=device) + value = storage[_OUTPUT_GUARD:_OUTPUT_GUARD + batch_size] + value.fill_(interior_value) + return GuardedOutput(storage=storage, value=value, guard_value=guard_value) + + +def _optional_sampling_outputs( + batch_size: int, + device: torch.device | str, +) -> tuple[GuardedOutput, GuardedOutput, GuardedOutput]: + sampled_logprobs = _guarded_output(batch_size * _MAX_LOGPROB, + torch.float32, device, -91.0) + sampled_indexes = _guarded_output(batch_size * _MAX_LOGPROB, torch.int32, + device, -193) + sampled_nums = _guarded_output(batch_size, torch.int32, device, -307) + return sampled_logprobs, sampled_indexes, sampled_nums + + +def _strided_rows( + batch_size: int, + width: int, + padding: int, + storage_offset: int, + dtype: torch.dtype, + device: torch.device | str, + fill_value: float | int = 0) -> tuple[torch.Tensor, torch.Tensor]: + stride = width + padding + storage_size = storage_offset + (batch_size - 1) * stride + width + 7 + storage = torch.full((storage_size, ), + fill_value, + dtype=dtype, + device=device) + view = torch.as_strided(storage, (batch_size, width), (stride, 1), + storage_offset=storage_offset) + return storage, view + + +def _fixed_seeds(count: int, + device: torch.device | str, + start: int = 1729) -> torch.Tensor: + return torch.arange(start, start + count, dtype=torch.int64, + device=device).to(torch.uint64) + + +def _initialize_state_storage(random_state_count: int, + device: torch.device | str, + start_seed: int = 1729) -> torch.Tensor: + random_states = allocate_random_states(random_state_count, device) + initialize_random_states( + random_states, _fixed_seeds(random_state_count, device, start_seed), + torch.ones(random_state_count, dtype=torch.bool, device=device)) + return random_states + + +def test_processed_sampling_null_mask_matches_all_true_with_optional_outputs( +) -> None: + device = torch.device('cuda') + batch_size = 4 + width = 3 + stream = torch.cuda.Stream(device=device) + + with torch.cuda.stream(stream): + probability_storage, probabilities = _strided_rows( + batch_size, + width, + padding=7, + storage_offset=5, + dtype=torch.float32, + device=device, + fill_value=-17.0, + ) + index_storage, indices = _strided_rows( + batch_size, + width, + padding=7, + storage_offset=3, + dtype=torch.int32, + device=device, + fill_value=-29, + ) + probabilities.zero_() + probabilities[:, 0] = 1.0 + indices.copy_( + torch.arange(101, + 101 + batch_size * width, + dtype=torch.int32, + device=device).view(batch_size, width)) + kept = torch.ones(batch_size, dtype=torch.int32, device=device) + state_indices = torch.tensor([4, 1, 3, 0], + dtype=torch.int32, + device=device) + + first_states = allocate_random_states(5, device) + second_states = allocate_random_states(5, device) + seeds = _fixed_seeds(5, device, start=31001) + initialize = torch.ones(5, dtype=torch.bool, device=device) + initialize_random_states(first_states, seeds, initialize) + initialize_random_states(second_states, seeds, initialize) + + first_selected = _guarded_output(batch_size, torch.int32, device, -777) + second_selected = _guarded_output(batch_size, torch.int32, device, + -777) + first_logprobs, first_indexes, first_nums = ( + _optional_sampling_outputs(batch_size, device)) + second_logprobs, second_indexes, second_nums = ( + _optional_sampling_outputs(batch_size, device)) + all_true = torch.ones(batch_size, dtype=torch.bool, device=device) + probability_snapshot = probability_storage.clone() + index_snapshot = index_storage.clone() + + sample_processed_probabilities( + probabilities, + indices, + kept, + first_states, + state_indices, + None, + first_selected.value, + first_logprobs.value.view(batch_size, _MAX_LOGPROB), + first_indexes.value.view(batch_size, _MAX_LOGPROB), + first_nums.value, + ) + sample_processed_probabilities( + probabilities, + indices, + kept, + second_states, + state_indices, + all_true, + second_selected.value, + second_logprobs.value.view(batch_size, _MAX_LOGPROB), + second_indexes.value.view(batch_size, _MAX_LOGPROB), + second_nums.value, + ) + + stream.synchronize() + + assert probabilities.stride(0) == indices.stride(0) == width + 7 + assert probabilities.storage_offset() == 5 + assert indices.storage_offset() == 3 + assert torch.equal(first_selected.value, indices[:, 0]) + assert torch.equal(first_selected.storage, second_selected.storage) + assert torch.equal(first_logprobs.storage, second_logprobs.storage) + assert torch.equal(first_indexes.storage, second_indexes.storage) + assert torch.equal(first_nums.storage, second_nums.storage) + assert torch.equal(first_states, second_states) + assert torch.equal(first_nums.value, torch.ones_like(first_nums.value)) + assert torch.equal( + first_indexes.value.view(batch_size, _MAX_LOGPROB)[:, 0], indices[:, + 0]) + assert torch.equal( + first_logprobs.value.view(batch_size, _MAX_LOGPROB)[:, 0], + torch.zeros(batch_size, dtype=torch.float32, device=device)) + assert torch.equal(probability_storage, probability_snapshot) + assert torch.equal(index_storage, index_snapshot) + for output in (first_selected, second_selected, first_logprobs, + second_logprobs, first_indexes, second_indexes, first_nums, + second_nums): + output.assert_guards() + + +def test_processed_sampling_all_false_mask_ignores_poisoned_rows() -> None: + device = torch.device('cuda') + batch_size = 3 + probabilities = torch.full((batch_size, 4), + -12345.0, + dtype=torch.float32, + device=device) + indices = torch.full((batch_size, 4), + -23456, + dtype=torch.int32, + device=device) + kept = torch.tensor([-101, -202, -303], dtype=torch.int32, device=device) + state_indices = torch.tensor([-401, -502, -603], + dtype=torch.int32, + device=device) + sample_mask = torch.zeros(batch_size, dtype=torch.bool, device=device) + selected = _guarded_output(batch_size, torch.int32, device, -777) + sampled_logprobs, sampled_indexes, sampled_nums = ( + _optional_sampling_outputs(batch_size, device)) + probability_snapshot = probabilities.clone() + index_snapshot = indices.clone() + + sample_processed_probabilities( + probabilities, + indices, + kept, + torch.empty(0, dtype=torch.uint8, device=device), + state_indices, + sample_mask, + selected.value, + sampled_logprobs.value.view(batch_size, _MAX_LOGPROB), + sampled_indexes.value.view(batch_size, _MAX_LOGPROB), + sampled_nums.value, + ) + torch.cuda.current_stream(device).synchronize() + + assert torch.equal(selected.value, torch.full_like(selected.value, -777)) + assert torch.equal(sampled_logprobs.value, + torch.full_like(sampled_logprobs.value, -91.0)) + assert torch.equal(sampled_indexes.value, + torch.full_like(sampled_indexes.value, -193)) + assert torch.equal(sampled_nums.value, + torch.full_like(sampled_nums.value, -307)) + assert torch.equal(probabilities, probability_snapshot) + assert torch.equal(indices, index_snapshot) + selected.assert_guards() + sampled_logprobs.assert_guards() + sampled_indexes.assert_guards() + sampled_nums.assert_guards() + + +def test_processed_sampling_mixed_mask_rng_and_all_null_outputs() -> None: + device = torch.device('cuda') + batch_size = 5 + width = 4 + sample_mask = torch.tensor([True, False, True, False, True], + dtype=torch.bool, + device=device) + probabilities = torch.full((batch_size, width), + -73.0, + dtype=torch.float32, + device=device) + indices = torch.full((batch_size, width), + -89, + dtype=torch.int32, + device=device) + kept = torch.tensor([1, -103, 1, -107, 1], + dtype=torch.int32, + device=device) + selected_values = torch.tensor([501, -1, 701, -1, 901], + dtype=torch.int32, + device=device) + probabilities[sample_mask, 0] = 1.0 + indices[:, 0] = selected_values + state_indices = torch.tensor([5, 3, 1, -109, 4], + dtype=torch.int32, + device=device) + + mixed_states = allocate_random_states(6, device) + active_only_states = allocate_random_states(6, device) + null_output_states = allocate_random_states(6, device) + seeds = _fixed_seeds(6, device, start=33001) + initialize = torch.ones(6, dtype=torch.bool, device=device) + initialize_random_states(mixed_states, seeds, initialize) + initialize_random_states(active_only_states, seeds, initialize) + initialize_random_states(null_output_states, seeds, initialize) + initial_states = mixed_states.clone() + + mixed_selected = _guarded_output(batch_size, torch.int32, device, -777) + null_output_selected = _guarded_output(batch_size, torch.int32, device, + -777) + active_only_selected = _guarded_output(3, torch.int32, device, -777) + sampled_logprobs, sampled_indexes, sampled_nums = ( + _optional_sampling_outputs(batch_size, device)) + probability_snapshot = probabilities.clone() + index_snapshot = indices.clone() + active_probabilities = probabilities[sample_mask].contiguous() + active_indices = indices[sample_mask].contiguous() + active_kept = kept[sample_mask].contiguous() + active_state_indices = state_indices[sample_mask].contiguous() + + sample_processed_probabilities( + probabilities, + indices, + kept, + mixed_states, + state_indices, + sample_mask, + mixed_selected.value, + sampled_logprobs.value.view(batch_size, _MAX_LOGPROB), + sampled_indexes.value.view(batch_size, _MAX_LOGPROB), + sampled_nums.value, + ) + sample_processed_probabilities( + probabilities, + indices, + kept, + null_output_states, + state_indices, + sample_mask, + null_output_selected.value, + None, + None, + None, + ) + sample_processed_probabilities( + active_probabilities, + active_indices, + active_kept, + active_only_states, + active_state_indices, + None, + active_only_selected.value, + None, + None, + None, + ) + torch.cuda.current_stream(device).synchronize() + + expected_active = selected_values[sample_mask] + assert torch.equal(mixed_selected.value[sample_mask], expected_active) + assert torch.equal(null_output_selected.value[sample_mask], + expected_active) + assert torch.equal(active_only_selected.value, expected_active) + assert torch.equal( + mixed_selected.value[~sample_mask], + torch.full((2, ), -777, dtype=torch.int32, device=device)) + assert torch.equal( + null_output_selected.value[~sample_mask], + torch.full((2, ), -777, dtype=torch.int32, device=device)) + assert torch.equal(mixed_states, active_only_states) + assert torch.equal(mixed_states, null_output_states) + assert not torch.equal(mixed_states, initial_states) + + logprob_rows = sampled_logprobs.value.view(batch_size, _MAX_LOGPROB) + index_rows = sampled_indexes.value.view(batch_size, _MAX_LOGPROB) + assert torch.equal(logprob_rows[sample_mask, 0], + torch.zeros(3, dtype=torch.float32, device=device)) + assert torch.equal(index_rows[sample_mask, 0], expected_active) + assert torch.equal( + logprob_rows[~sample_mask], + torch.full((2, _MAX_LOGPROB), + -91.0, + dtype=torch.float32, + device=device)) + assert torch.equal( + index_rows[~sample_mask], + torch.full((2, _MAX_LOGPROB), -193, dtype=torch.int32, device=device)) + assert torch.equal(sampled_nums.value[sample_mask], + torch.ones(3, dtype=torch.int32, device=device)) + assert torch.equal( + sampled_nums.value[~sample_mask], + torch.full((2, ), -307, dtype=torch.int32, device=device)) + assert torch.equal(probabilities, probability_snapshot) + assert torch.equal(indices, index_snapshot) + mixed_selected.assert_guards() + null_output_selected.assert_guards() + active_only_selected.assert_guards() + sampled_logprobs.assert_guards() + sampled_indexes.assert_guards() + sampled_nums.assert_guards() + + +def test_processed_sampling_and_append_empty_views() -> None: + device = torch.device('cuda') + width = 7 + probability_storage = torch.empty(width + 11, + dtype=torch.float32, + device=device) + index_storage = torch.empty(width + 11, dtype=torch.int32, device=device) + probabilities = torch.as_strided(probability_storage, (0, width), + (width + 3, 1), + storage_offset=4) + indices = torch.as_strided(index_storage, (0, width), (width + 3, 1), + storage_offset=2) + empty_int = torch.empty(0, dtype=torch.int32, device=device) + empty_bool = torch.empty(0, dtype=torch.bool, device=device) + + sample_processed_probabilities( + probabilities, + indices, + empty_int, + torch.empty(0, dtype=torch.uint8, device=device), + empty_int.clone(), + empty_bool, + empty_int.clone(), + torch.empty((0, _MAX_LOGPROB), dtype=torch.float32, device=device), + torch.empty((0, _MAX_LOGPROB), dtype=torch.int32, device=device), + empty_int.clone(), + ) + append_one_token_and_advance_sequence( + torch.empty(0, dtype=torch.int64, device=device), + empty_int.clone(), + empty_int.clone(), + ) + torch.cuda.current_stream(device).synchronize() + + +def test_append_one_token_uses_pointer_order_and_distinct_lengths() -> None: + device = torch.device('cuda') + batch_size = 4 + session_length = 10 + row_guard = 6 + row_guard_value = -123456789 + row_initial_value = -41 + allocations = [ + torch.full((session_length + 2 * row_guard, ), + row_guard_value, + dtype=torch.int32, + device=device) for _ in range(batch_size) + ] + rows = [] + for storage in allocations: + row = storage[row_guard:row_guard + session_length] + row.fill_(row_initial_value) + rows.append(row) + + pointer_order = [2, 0, 3, 1] + token_ids_ptrs = torch.tensor( + [rows[index].data_ptr() for index in pointer_order], + dtype=torch.int64, + device=device) + selected_values = [1101, 1102, 1103, 1104] + selected_tokens = torch.tensor(selected_values, + dtype=torch.int32, + device=device) + starting_lengths = [0, 3, 1, 5] + sequence_length = torch.tensor(starting_lengths, + dtype=torch.int32, + device=device) + allocation_snapshots = [storage.clone() for storage in allocations] + + append_one_token_and_advance_sequence(token_ids_ptrs, selected_tokens, + sequence_length) + torch.cuda.current_stream(device).synchronize() + + assert torch.equal( + sequence_length, + torch.tensor([value + 1 for value in starting_lengths], + dtype=torch.int32, + device=device)) + for logical_row, physical_row in enumerate(pointer_order): + expected = allocation_snapshots[physical_row].clone() + expected[row_guard + + starting_lengths[logical_row]] = selected_values[logical_row] + assert torch.equal(allocations[physical_row], expected) + assert torch.equal( + allocations[physical_row][:row_guard], + torch.full((row_guard, ), + row_guard_value, + dtype=torch.int32, + device=device)) + assert torch.equal( + allocations[physical_row][-row_guard:], + torch.full((row_guard, ), + row_guard_value, + dtype=torch.int32, + device=device)) + + +def test_append_successive_launches_preserve_nondefault_stream_order() -> None: + device = torch.device('cuda') + batch_size = 2 + session_length = 8 + row_guard = 4 + row_guard_value = -987654321 + row_initial_value = -53 + allocations = [ + torch.full((session_length + 2 * row_guard, ), + row_guard_value, + dtype=torch.int32, + device=device) for _ in range(batch_size) + ] + rows = [] + for storage in allocations: + row = storage[row_guard:row_guard + session_length] + row.fill_(row_initial_value) + rows.append(row) + + pointer_order = [1, 0] + logical_rows = [rows[index] for index in pointer_order] + token_ids_ptrs = torch.tensor([row.data_ptr() for row in logical_rows], + dtype=torch.int64, + device=device) + starting_lengths = [1, 3] + sequence_length = torch.tensor(starting_lengths, + dtype=torch.int32, + device=device) + first_values = [2101, 2102] + second_values = [2201, 2202] + stream = torch.cuda.Stream(device=device) + torch.cuda.current_stream(device).synchronize() + + with torch.cuda.stream(stream): + first_selected_tokens = torch.tensor(first_values, + dtype=torch.int32, + device=device) + second_selected_tokens = torch.tensor(second_values, + dtype=torch.int32, + device=device) + append_one_token_and_advance_sequence( + token_ids_ptrs, + first_selected_tokens, + sequence_length, + ) + first_length_snapshot = sequence_length.clone() + first_row_snapshot = torch.stack(logical_rows) + append_one_token_and_advance_sequence( + token_ids_ptrs, + second_selected_tokens, + sequence_length, + ) + final_length_snapshot = sequence_length.clone() + final_row_snapshot = torch.stack(logical_rows) + + stream.synchronize() + + expected_first_rows = torch.full((batch_size, session_length), + row_initial_value, + dtype=torch.int32, + device=device) + expected_final_rows = expected_first_rows.clone() + for row in range(batch_size): + expected_first_rows[row, starting_lengths[row]] = first_values[row] + expected_final_rows[row, starting_lengths[row]] = first_values[row] + expected_final_rows[row, + starting_lengths[row] + 1] = second_values[row] + assert first_length_snapshot.device.type == 'cuda' + assert torch.equal( + first_length_snapshot, + torch.tensor([value + 1 for value in starting_lengths], + dtype=torch.int32, + device=device)) + assert torch.equal(first_row_snapshot, expected_first_rows) + assert torch.equal( + final_length_snapshot, + torch.tensor([value + 2 for value in starting_lengths], + dtype=torch.int32, + device=device)) + assert torch.equal(final_row_snapshot, expected_final_rows) + assert torch.equal(sequence_length, final_length_snapshot) + for storage in allocations: + assert torch.equal( + storage[:row_guard], + torch.full((row_guard, ), + row_guard_value, + dtype=torch.int32, + device=device)) + assert torch.equal( + storage[-row_guard:], + torch.full((row_guard, ), + row_guard_value, + dtype=torch.int32, + device=device)) + + +def _block_token_rows( + request_count: int, + device: torch.device, + row_width: int = 24, +) -> tuple[list[torch.Tensor], list[torch.Tensor], torch.Tensor]: + guard = 5 + guard_value = -123456789 + allocations = [ + torch.full((row_width + 2 * guard, ), + guard_value, + dtype=torch.int32, + device=device) for _ in range(request_count) + ] + rows = [storage[guard:guard + row_width] for storage in allocations] + for row in rows: + row.fill_(-41) + pointers = torch.tensor([row.data_ptr() for row in rows], + dtype=torch.int64, + device=device) + return allocations, rows, pointers + + +def _assert_block_row_guards(allocations: list[torch.Tensor]) -> None: + for storage in allocations: + assert torch.equal(storage[:5], + torch.full_like(storage[:5], -123456789)) + assert torch.equal(storage[-5:], + torch.full_like(storage[-5:], -123456789)) + + +def test_verify_target_block_mixed_rows_padded_position_major_inputs() -> None: + device = torch.device('cuda') + request_count = 5 + generation_count = 4 + position_count = 4 + rows_count = generation_count * position_count + width = 3 + + _, probabilities = _strided_rows(rows_count, + width, + padding=7, + storage_offset=3, + dtype=torch.float32, + device=device) + _, token_ids = _strided_rows(rows_count, + width, + padding=11, + storage_offset=5, + dtype=torch.int32, + device=device) + probabilities.zero_() + token_ids.fill_(-1) + + def point(position: int, generation: int, token: int) -> None: + flat = position * generation_count + generation + probabilities[flat, 0] = 1.0 + token_ids[flat, 0] = token + + point(0, 0, 101) + for position, token in enumerate((201, 202, 203, 204)): + point(position, 1, token) + probabilities[2, :2] = torch.tensor([0.0, 1.0], device=device) + token_ids[2, :2] = torch.tensor([301, 399], + dtype=torch.int32, + device=device) + for position, token in enumerate((401, 498, 403, 404)): + point(position, 3, token) + + kept = torch.ones(rows_count, dtype=torch.int32, device=device) + kept[2] = 2 + draft_storage = torch.full( + ((position_count - 1) * (generation_count + 2), ), + -17, + dtype=torch.int32, + device=device) + draft_ids = torch.as_strided(draft_storage, + (position_count - 1, generation_count), + (generation_count + 2, 1)) + draft_ids.copy_( + torch.tensor( + [[0, 201, 301, 401], [0, 202, 302, 402], [0, 203, 303, 403]], + dtype=torch.int32, + device=device)) + greedy = torch.ones(rows_count, dtype=torch.bool, device=device) + greedy[2] = False + logits_active = torch.ones(rows_count, dtype=torch.bool, device=device) + random_states = _initialize_state_storage(generation_count, device, 41001) + state_indices = torch.arange(generation_count, + dtype=torch.int32, + device=device) + allocations, token_rows, token_ptrs = _block_token_rows( + request_count, device) + entry_lengths = torch.tensor([2, 3, 5, 1, 4], + dtype=torch.int32, + device=device) + offsets = torch.tensor([0, 1, 2, 2, 3, 4], + dtype=torch.int32, + device=device) + speculative = torch.tensor([False, True, False, True, True], + dtype=torch.bool, + device=device) + selected_storage = torch.full((request_count * (position_count + 2), ), + -777, + dtype=torch.int32, + device=device) + selected = torch.as_strided(selected_storage, + (request_count, position_count), + (position_count + 2, 1)) + accept_len = torch.full((request_count, ), + -9, + dtype=torch.int32, + device=device) + accepted_count = torch.full((request_count, ), + -7, + dtype=torch.int32, + device=device) + + verify_target_block(probabilities, token_ids, kept, draft_ids, greedy, + logits_active, random_states, state_indices, + token_ptrs, entry_lengths, offsets, speculative, + selected, accept_len, accepted_count, position_count) + torch.cuda.current_stream(device).synchronize() + + assert accept_len.tolist() == [1, 4, -9, 1, 2] + assert selected[0, 0].item() == 101 + assert selected[1].tolist() == [201, 202, 203, 204] + assert selected[3, :2].tolist() == [399, -777] + assert selected[4, :2].tolist() == [401, 498] + assert accepted_count.tolist() == [-7, 3, -7, 0, 1] + assert token_rows[0][2].item() == 101 + assert token_rows[1][3:7].tolist() == [201, 202, 203, 204] + assert token_rows[2].eq(-41).all() + assert token_rows[3][1].item() == 399 + assert token_rows[4][4:6].tolist() == [401, 498] + assert selected_storage.view(request_count, position_count + + 2)[:, position_count:].eq(-777).all() + _assert_block_row_guards(allocations) + + +@pytest.mark.parametrize('rejection_position', [0, 1, 2]) +def test_verify_target_block_rejection_position_and_rng_advance( + rejection_position: int) -> None: + device = torch.device('cuda') + position_count = 4 + probabilities = torch.zeros((position_count, 2), + dtype=torch.float32, + device=device) + token_ids = torch.empty((position_count, 2), + dtype=torch.int32, + device=device) + drafts = torch.tensor([[501], [502], [503]], + dtype=torch.int32, + device=device) + for position in range(position_count - 1): + token_ids[position] = torch.tensor([501 + position, 601 + position], + dtype=torch.int32, + device=device) + if position == rejection_position: + probabilities[position] = torch.tensor([0.0, 1.0], device=device) + else: + probabilities[position, 0] = 1.0 + probabilities[-1, 0] = 1.0 + token_ids[-1] = torch.tensor([700, 701], dtype=torch.int32, device=device) + + block_states = _initialize_state_storage(1, device, 42001) + control_states = _initialize_state_storage(1, device, 42001) + _, _, token_ptrs = _block_token_rows(1, device) + selected = torch.full((1, position_count), + -777, + dtype=torch.int32, + device=device) + accept_len = torch.full((1, ), -1, dtype=torch.int32, device=device) + accepted_count = torch.full((1, ), -1, dtype=torch.int32, device=device) + offsets = torch.tensor([0, 1], dtype=torch.int32, device=device) + state_indices = torch.zeros(1, dtype=torch.int32, device=device) + + verify_target_block( + probabilities, + token_ids, + torch.tensor([2] * position_count, dtype=torch.int32, device=device), + drafts, + torch.zeros(position_count, dtype=torch.bool, device=device), + torch.ones(position_count, dtype=torch.bool, device=device), + block_states, + state_indices, + token_ptrs, + torch.zeros(1, dtype=torch.int32, device=device), + offsets, + torch.ones(1, dtype=torch.bool, device=device), + selected, + accept_len, + accepted_count, + position_count, + ) + block_accept_len = accept_len.clone() + + ordinary_probabilities = torch.tensor([[1.0]], + dtype=torch.float32, + device=device) + ordinary_ids = torch.tensor([[900]], dtype=torch.int32, device=device) + empty_drafts = torch.empty((0, 1), dtype=torch.int32, device=device) + for _ in range(rejection_position + 2): + verify_target_block( + ordinary_probabilities, + ordinary_ids, + torch.ones(1, dtype=torch.int32, device=device), + empty_drafts, + torch.zeros(1, dtype=torch.bool, device=device), + torch.ones(1, dtype=torch.bool, device=device), + control_states, + state_indices, + token_ptrs, + torch.zeros(1, dtype=torch.int32, device=device), + offsets, + torch.zeros(1, dtype=torch.bool, device=device), + selected[:, :1], + accept_len, + None, + 1, + ) + torch.cuda.current_stream(device).synchronize() + + assert block_accept_len.item() == rejection_position + 1 + assert accepted_count.item() == rejection_position + assert torch.equal(block_states, control_states) + + +def test_verify_target_block_full_acceptance_rng_and_null_metrics() -> None: + device = torch.device('cuda') + position_count = 4 + probabilities = torch.zeros((position_count, 2), + dtype=torch.float32, + device=device) + probabilities[:, 0] = 1.0 + token_ids = torch.tensor([[801, 901], [802, 902], [803, 903], [804, 904]], + dtype=torch.int32, + device=device) + drafts = torch.tensor([[801], [802], [803]], + dtype=torch.int32, + device=device) + block_states = _initialize_state_storage(1, device, 43001) + control_states = _initialize_state_storage(1, device, 43001) + allocations, token_rows, token_ptrs = _block_token_rows(1, device) + selected = torch.full((1, position_count), + -777, + dtype=torch.int32, + device=device) + accept_len = torch.zeros(1, dtype=torch.int32, device=device) + offsets = torch.tensor([0, 1], dtype=torch.int32, device=device) + state_indices = torch.zeros(1, dtype=torch.int32, device=device) + + common = dict( + probabilities=probabilities, + probability_token_ids=token_ids, + kept_count=torch.ones(position_count, dtype=torch.int32, + device=device), + verification_draft_ids=drafts, + greedy=torch.zeros(position_count, dtype=torch.bool, device=device), + logits_active=torch.ones(position_count, + dtype=torch.bool, + device=device), + random_states=block_states, + random_state_indices=state_indices, + request_token_ids_ptrs=token_ptrs, + entry_sequence_length=torch.tensor([2], + dtype=torch.int32, + device=device), + request_to_generation_offsets=offsets, + speculative_row=torch.ones(1, dtype=torch.bool, device=device), + selected_span_ids=selected, + accept_len=accept_len, + accepted_draft_count=None, + position_count=position_count, + ) + verify_target_block(**common) + + ordinary_probabilities = torch.tensor([[1.0]], + dtype=torch.float32, + device=device) + ordinary_ids = torch.tensor([[999]], dtype=torch.int32, device=device) + empty_drafts = torch.empty((0, 1), dtype=torch.int32, device=device) + for _ in range(position_count): + verify_target_block(ordinary_probabilities, ordinary_ids, + torch.ones(1, dtype=torch.int32, + device=device), empty_drafts, + torch.zeros(1, dtype=torch.bool, device=device), + torch.ones(1, dtype=torch.bool, device=device), + control_states, state_indices, token_ptrs, + torch.zeros(1, dtype=torch.int32, + device=device), offsets, + torch.zeros(1, dtype=torch.bool, device=device), + selected[:, :1], accept_len, None, 1) + torch.cuda.current_stream(device).synchronize() + + assert torch.equal(block_states, control_states) + assert token_rows[0][2:6].tolist() == [801, 802, 803, 804] + _assert_block_row_guards(allocations) diff --git a/tests/turbomind/speculative_sequence/__init__.py b/tests/turbomind/speculative_sequence/__init__.py new file mode 100644 index 0000000000..ef101fec61 --- /dev/null +++ b/tests/turbomind/speculative_sequence/__init__.py @@ -0,0 +1 @@ +# Copyright (c) OpenMMLab. All rights reserved. diff --git a/tests/turbomind/speculative_sequence/reference.py b/tests/turbomind/speculative_sequence/reference.py new file mode 100644 index 0000000000..96e1b87385 --- /dev/null +++ b/tests/turbomind/speculative_sequence/reference.py @@ -0,0 +1,273 @@ +# Copyright (c) OpenMMLab. All rights reserved. + +import torch + + +def initialize_target_verification_reference( + token_rows: list[torch.Tensor | None], + entry_sequence_length: torch.Tensor, + finished_on_entry: torch.Tensor, + speculative_row: torch.Tensor, + request_to_generation_row_offsets: torch.Tensor, + draft_count: int, + metrics_enabled: bool, +): + offsets = request_to_generation_row_offsets.cpu().tolist() + entry_lengths = entry_sequence_length.cpu().tolist() + inherited_finished = finished_on_entry.cpu().tolist() + speculative = speculative_row.cpu().tolist() + generation_count = offsets[-1] + + block_logits_active = torch.empty( + (draft_count + 1, generation_count), + dtype=torch.bool, + device=entry_sequence_length.device, + ) + effective_history = torch.empty( + (draft_count + 1, generation_count), + dtype=torch.int32, + device=entry_sequence_length.device, + ) + verification_draft_ids = torch.empty( + (draft_count, generation_count), + dtype=torch.int32, + device=entry_sequence_length.device, + ) + accepted_draft_count = (torch.empty( + len(token_rows), + dtype=torch.int32, + device=entry_sequence_length.device, + ) if metrics_enabled else None) + + for b, row in enumerate(token_rows): + active = not inherited_finished[b] + verify = active and speculative[b] + if accepted_draft_count is not None: + accepted_draft_count[b] = 0 if verify else -1 + + if offsets[b + 1] == offsets[b]: + continue + + g = offsets[b] + S = entry_lengths[b] + for i in range(draft_count + 1): + block_logits_active[i, g] = active and (i == 0 or verify) + effective_history[i, g] = S + i + for i in range(draft_count): + verification_draft_ids[i, g] = row[S + i] if verify else 0 + + return ( + block_logits_active, + effective_history, + verification_draft_ids, + accepted_draft_count, + ) + + +def build_draft_extension_key_offsets_reference( + q_offsets: torch.Tensor, + entry_sequence_length: torch.Tensor, + accept_len: torch.Tensor, + extension_index: int, +) -> torch.Tensor: + q_width = q_offsets[1:] - q_offsets[:-1] + key_len = torch.where( + q_width == 1, + entry_sequence_length + accept_len + extension_index, + 0, + ) + return torch.cat([ + torch.zeros(1, dtype=torch.int32, device=key_len.device), + torch.cumsum(key_len, dim=0, dtype=torch.int32), + ]) + + +def pack_stop_words( + phrases_by_row: list[list[list[int]]], + width: int, + device: torch.device | str, +) -> torch.Tensor: + packed = torch.full( + (len(phrases_by_row), 2, width), + -1, + dtype=torch.int32, + device=device, + ) + for b, phrases in enumerate(phrases_by_row): + flat_words = [] + ends = [] + for phrase in phrases: + flat_words.extend(phrase) + ends.append(len(flat_words)) + if flat_words: + packed[b, 0, :len(flat_words)] = torch.tensor(flat_words, + dtype=torch.int32, + device=device) + if ends: + packed[b, 1, :len(ends)] = torch.tensor(ends, + dtype=torch.int32, + device=device) + return packed + + +def stop_criteria_reference( + token_rows: list[torch.Tensor], + sequence_length: torch.Tensor, + stop_words: torch.Tensor | None, + sequence_length_limit: torch.Tensor, + finished: torch.Tensor, + accept_len: torch.Tensor | None = None, +): + output_finished = finished.clone() + output_accept_len = None if accept_len is None else accept_len.clone() + lengths = sequence_length.cpu().tolist() + limits = sequence_length_limit.cpu().tolist() + + packed = None if stop_words is None else stop_words.cpu().tolist() + + for b, token_row in enumerate(token_rows): + if output_accept_len is not None: + if bool(output_finished[b].item()): + output_accept_len[b] = 0 + continue + entry_len = lengths[b] + span_len = int(output_accept_len[b].item()) + else: + if bool(output_finished[b].item()): + continue + current_len = lengths[b] + if current_len <= 0: + continue + entry_len = current_len - 1 + span_len = 1 + + if span_len <= 0: + continue + + tokens = token_row.cpu().tolist() + phrases = [] + if packed is not None: + phrase_begin = 0 + for phrase_end in packed[b][1]: + if phrase_end < 0: + break + phrases.append(packed[b][0][phrase_begin:phrase_end]) + phrase_begin = phrase_end + + for j in range(span_len): + effective_len = entry_len + j + 1 + terminal = effective_len >= limits[b] + + if not terminal: + for phrase in phrases: + phrase_size = len(phrase) + if phrase_size <= 0 or effective_len < phrase_size: + continue + if tokens[effective_len - + phrase_size:effective_len] == phrase: + terminal = True + break + + if terminal: + if output_accept_len is not None: + output_accept_len[b] = j + 1 + output_finished[b] = True + break + + return output_accept_len, output_finished + + +def build_draft_refresh_inputs_reference( + token_rows: list[torch.Tensor | None], + refresh_q_offsets: torch.Tensor, + refresh_k_offsets: torch.Tensor, + extension_q_offsets: torch.Tensor, + accept_len: torch.Tensor, + limit_to_accept_len: torch.Tensor, + finished: torch.Tensor, +): + q_offsets = refresh_q_offsets.cpu().tolist() + k_offsets = refresh_k_offsets.cpu().tolist() + extension_offsets = extension_q_offsets.cpu().tolist() + limiting = limit_to_accept_len.cpu().tolist() + finished_rows = finished.cpu().tolist() + + draft_input_ids = torch.zeros( + q_offsets[-1], + dtype=torch.int32, + device=refresh_q_offsets.device, + ) + selected_token_pos = torch.zeros( + extension_offsets[-1], + dtype=torch.int32, + device=refresh_q_offsets.device, + ) + candidate_active = torch.zeros( + extension_offsets[-1], + dtype=torch.bool, + device=refresh_q_offsets.device, + ) + + for b, token_row in enumerate(token_rows): + q_begin = q_offsets[b] + q_end = q_offsets[b + 1] + q_len = q_end - q_begin + + if extension_offsets[b + 1] != extension_offsets[b]: + candidate = extension_offsets[b] + if not finished_rows[b]: + if limiting[b]: + committed = int(accept_len[b].item()) + if committed > 0: + candidate_active[candidate] = True + selected_token_pos[candidate] = q_begin + committed - 1 + elif q_len > 0: + candidate_active[candidate] = True + selected_token_pos[candidate] = q_end - 1 + + if q_len == 0: + continue + + token_begin = k_offsets[b + 1] - k_offsets[b] - q_len + 1 + valid_len = int(accept_len[b].item()) if limiting[b] else q_len + for j in range(valid_len): + draft_input_ids[q_begin + j] = token_row[token_begin + j] + + return draft_input_ids, selected_token_pos, candidate_active + + +def draft_argmax_and_store_token_reference( + logits: torch.Tensor, + token_rows: list[torch.Tensor | None], + extension_q_offsets: torch.Tensor, + candidate_active: torch.Tensor, + entry_sequence_length: torch.Tensor, + accept_len: torch.Tensor, + proposal_index: int, + vocab_size: int, +): + output_rows = [None if row is None else row.clone() for row in token_rows] + proposal_ids = torch.zeros( + logits.shape[0], + dtype=torch.int32, + device=logits.device, + ) + extension_offsets = extension_q_offsets.cpu().tolist() + + for b, output_row in enumerate(output_rows): + if extension_offsets[b + 1] == extension_offsets[b]: + continue + + candidate = extension_offsets[b] + if not bool(candidate_active[candidate].item()): + continue + + row = logits[candidate, :vocab_size].float() + row = torch.where(torch.isnan(row), -torch.inf, row) + token_id = torch.argmax(row).to(torch.int32) + proposal_ids[candidate] = token_id + token_position = (int(entry_sequence_length[b].item()) + + int(accept_len[b].item()) + proposal_index) + output_row[token_position] = token_id + + return output_rows, proposal_ids diff --git a/tests/turbomind/speculative_sequence/speculative_sequence.py b/tests/turbomind/speculative_sequence/speculative_sequence.py new file mode 100644 index 0000000000..c1f9da48b5 --- /dev/null +++ b/tests/turbomind/speculative_sequence/speculative_sequence.py @@ -0,0 +1,157 @@ +# Copyright (c) OpenMMLab. All rights reserved. + +import torch + + +def initialize_target_verification( + block_logits_active: torch.Tensor, + effective_history: torch.Tensor, + verification_draft_ids: torch.Tensor, + request_token_ids_ptrs: torch.Tensor, + entry_sequence_length: torch.Tensor, + finished_on_entry: torch.Tensor, + speculative_row: torch.Tensor, + accepted_draft_count: torch.Tensor | None, + request_to_generation_row_offsets: torch.Tensor, +) -> None: + import _turbomind + + stream = torch.cuda.current_stream(entry_sequence_length.device) + + _turbomind.initialize_target_verification( + block_logits_active, + effective_history, + verification_draft_ids, + request_token_ids_ptrs, + entry_sequence_length, + finished_on_entry, + speculative_row, + accepted_draft_count, + request_to_generation_row_offsets, + stream.cuda_stream, + ) + + +def build_draft_extension_key_offsets( + k_offsets: torch.Tensor, + q_offsets: torch.Tensor, + entry_sequence_length: torch.Tensor, + accept_len: torch.Tensor, + extension_index: int, +) -> None: + import _turbomind + + stream = torch.cuda.current_stream(q_offsets.device) + + _turbomind.build_draft_extension_key_offsets( + k_offsets, + q_offsets, + entry_sequence_length, + accept_len, + extension_index, + stream.cuda_stream, + ) + + +def stop_criteria( + token_ids_ptrs: torch.Tensor, + sequence_length: torch.Tensor, + stop_words: torch.Tensor | None, + sequence_length_limit: torch.Tensor, + finished: torch.Tensor, +) -> None: + import _turbomind + + stream = torch.cuda.current_stream(token_ids_ptrs.device) + + _turbomind.stop_criteria( + token_ids_ptrs, + sequence_length, + stop_words, + sequence_length_limit, + finished, + stream.cuda_stream, + ) + + +def speculative_stop_criteria( + token_ids_ptrs: torch.Tensor, + entry_sequence_length: torch.Tensor, + accept_len: torch.Tensor, + stop_words: torch.Tensor | None, + sequence_length_limit: torch.Tensor, + finished: torch.Tensor, +) -> None: + import _turbomind + + stream = torch.cuda.current_stream(token_ids_ptrs.device) + + _turbomind.speculative_stop_criteria( + token_ids_ptrs, + entry_sequence_length, + accept_len, + stop_words, + sequence_length_limit, + finished, + stream.cuda_stream, + ) + + +def build_draft_refresh_inputs( + draft_input_ids: torch.Tensor, + selected_token_pos: torch.Tensor, + candidate_active: torch.Tensor, + token_ids_ptrs: torch.Tensor, + refresh_q_offsets: torch.Tensor, + refresh_k_offsets: torch.Tensor, + extension_q_offsets: torch.Tensor, + accept_len: torch.Tensor, + limit_to_accept_len: torch.Tensor, + finished: torch.Tensor, +) -> None: + import _turbomind + + stream = torch.cuda.current_stream(refresh_q_offsets.device) + + _turbomind.build_draft_refresh_inputs( + draft_input_ids, + selected_token_pos, + candidate_active, + token_ids_ptrs, + refresh_q_offsets, + refresh_k_offsets, + extension_q_offsets, + accept_len, + limit_to_accept_len, + finished, + stream.cuda_stream, + ) + + +def draft_argmax_and_store_token( + logits: torch.Tensor, + proposal_ids: torch.Tensor, + token_ids_ptrs: torch.Tensor, + extension_q_offsets: torch.Tensor, + candidate_active: torch.Tensor, + entry_sequence_length: torch.Tensor, + accept_len: torch.Tensor, + proposal_index: int, + vocab_size: int, +) -> None: + import _turbomind + + stream = torch.cuda.current_stream(logits.device) + + _turbomind.draft_argmax_and_store_token( + logits, + proposal_ids, + token_ids_ptrs, + extension_q_offsets, + candidate_active, + entry_sequence_length, + accept_len, + proposal_index, + vocab_size, + stream.cuda_stream, + ) diff --git a/tests/turbomind/speculative_sequence/test_speculative_sequence.py b/tests/turbomind/speculative_sequence/test_speculative_sequence.py new file mode 100644 index 0000000000..bde7ceb224 --- /dev/null +++ b/tests/turbomind/speculative_sequence/test_speculative_sequence.py @@ -0,0 +1,1581 @@ +# Copyright (c) OpenMMLab. All rights reserved. + +import pytest +import torch + +from .reference import ( + build_draft_extension_key_offsets_reference, + build_draft_refresh_inputs_reference, + draft_argmax_and_store_token_reference, + initialize_target_verification_reference, + pack_stop_words, + stop_criteria_reference, +) +from .speculative_sequence import ( + build_draft_extension_key_offsets, + build_draft_refresh_inputs, + draft_argmax_and_store_token, + initialize_target_verification, + speculative_stop_criteria, + stop_criteria, +) + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), + reason='CUDA is required') + +_DEVICE = 'cuda' +_SPECULATIVE_K = 7 +_OUTPUT_SENTINEL = -7777777 +_TOKEN_SENTINEL = -6060606 +_POISON_SEQUENCE_LENGTH = 1700000000 +_POISON_ACCEPT_LEN = 900000000 + + +def _make_widths(batch_size: int, pattern: str) -> torch.Tensor: + if pattern == 'all_candidate': + return torch.ones(batch_size, dtype=torch.int32, device=_DEVICE) + if pattern == 'all_non_candidate': + return torch.zeros(batch_size, dtype=torch.int32, device=_DEVICE) + if pattern == 'mixed': + rows = torch.arange(batch_size, dtype=torch.int32, device=_DEVICE) + return ((rows % 3) != 1).to(torch.int32) + raise ValueError(f'unknown width pattern: {pattern}') + + +def _make_offsets(widths: torch.Tensor) -> torch.Tensor: + return torch.cat([ + torch.zeros(1, dtype=torch.int32, device=widths.device), + torch.cumsum(widths, dim=0, dtype=torch.int32), + ]) + + +def _make_case(batch_size: int, pattern: str): + widths = _make_widths(batch_size, pattern) + q_offsets = _make_offsets(widths) + rows = torch.arange(batch_size, dtype=torch.int32, device=_DEVICE) + entry_sequence_length = 9 + rows % 97 + accept_len = rows % 8 + + zero_width = widths == 0 + entry_sequence_length = torch.where( + zero_width, + torch.full_like(entry_sequence_length, _POISON_SEQUENCE_LENGTH), + entry_sequence_length, + ) + accept_len = torch.where( + zero_width, + torch.full_like(accept_len, _POISON_ACCEPT_LEN), + accept_len, + ) + return q_offsets, entry_sequence_length, accept_len + + +def _guarded_output(batch_size: int): + backing = torch.full((batch_size + 3, ), + _OUTPUT_SENTINEL, + dtype=torch.int32, + device=_DEVICE) + return backing, backing[1:-1] + + +def _make_token_rows(row_values): + backings = [] + rows = [] + for values in row_values: + backing = torch.full((len(values) + 2, ), + _TOKEN_SENTINEL, + dtype=torch.int32, + device=_DEVICE) + row = backing[1:-1] + row.copy_(torch.tensor(values, dtype=torch.int32, device=_DEVICE)) + backings.append(backing) + rows.append(row) + return backings, rows + + +def _token_pointer_array(rows): + return torch.tensor( + [0 if row is None else row.data_ptr() for row in rows], + dtype=torch.int64, + device=_DEVICE, + ) + + +def _assert_token_guards(backings): + for backing in backings: + assert backing[0].item() == _TOKEN_SENTINEL + assert backing[-1].item() == _TOKEN_SENTINEL + + +def _guarded_int_vector(values, sentinel=_OUTPUT_SENTINEL): + backing = torch.full((len(values) + 2, ), + sentinel, + dtype=torch.int32, + device=_DEVICE) + view = backing[1:-1] + if values: + view.copy_(torch.tensor(values, dtype=torch.int32, device=_DEVICE)) + return backing, view + + +def _guarded_bool_vector(values, sentinel=True): + backing = torch.full((len(values) + 2, ), + sentinel, + dtype=torch.bool, + device=_DEVICE) + view = backing[1:-1] + if values: + view.copy_(torch.tensor(values, dtype=torch.bool, device=_DEVICE)) + return backing, view + + +def _assert_vector_guards(backing, sentinel): + assert backing[0].item() == sentinel + assert backing[-1].item() == sentinel + + +def _stop_and_compare( + token_rows, + token_ids_ptrs, + sequence_length, + stop_words, + sequence_length_limit, + finished, + accept_len=None, +): + expected_accept_len, expected_finished = stop_criteria_reference( + token_rows, + sequence_length, + stop_words, + sequence_length_limit, + finished, + accept_len=accept_len, + ) + pointer_before = token_ids_ptrs.clone() + rows_before = [row.clone() for row in token_rows] + sequence_length_before = sequence_length.clone() + limit_before = sequence_length_limit.clone() + stop_words_before = None if stop_words is None else stop_words.clone() + + if accept_len is None: + result = stop_criteria( + token_ids_ptrs, + sequence_length, + stop_words, + sequence_length_limit, + finished, + ) + else: + result = speculative_stop_criteria( + token_ids_ptrs, + sequence_length, + accept_len, + stop_words, + sequence_length_limit, + finished, + ) + + assert result is None + assert torch.equal(finished, expected_finished) + if accept_len is not None: + assert torch.equal(accept_len, expected_accept_len) + assert torch.equal(token_ids_ptrs, pointer_before) + for actual, before in zip(token_rows, rows_before): + assert torch.equal(actual, before) + assert torch.equal(sequence_length, sequence_length_before) + assert torch.equal(sequence_length_limit, limit_before) + if stop_words is not None: + assert torch.equal(stop_words, stop_words_before) + + +def _refresh_and_compare( + token_rows, + token_ids_ptrs, + draft_input_ids, + selected_token_pos, + candidate_active, + refresh_q_offsets, + refresh_k_offsets, + extension_q_offsets, + accept_len, + limit_to_accept_len, + finished, +): + expected = build_draft_refresh_inputs_reference( + token_rows, + refresh_q_offsets, + refresh_k_offsets, + extension_q_offsets, + accept_len, + limit_to_accept_len, + finished, + ) + token_rows_before = [ + None if row is None else row.clone() for row in token_rows + ] + inputs = [ + token_ids_ptrs, + refresh_q_offsets, + refresh_k_offsets, + extension_q_offsets, + accept_len, + limit_to_accept_len, + finished, + ] + inputs_before = [value.clone() for value in inputs] + + result = build_draft_refresh_inputs( + draft_input_ids, + selected_token_pos, + candidate_active, + token_ids_ptrs, + refresh_q_offsets, + refresh_k_offsets, + extension_q_offsets, + accept_len, + limit_to_accept_len, + finished, + ) + + expected_ids, expected_pos, expected_active = expected + assert result is None + assert torch.equal(draft_input_ids, expected_ids) + assert torch.equal(selected_token_pos, expected_pos) + assert torch.equal(candidate_active, expected_active) + for actual, before in zip(inputs, inputs_before): + assert torch.equal(actual, before) + for actual, before in zip(token_rows, token_rows_before): + if actual is not None: + assert torch.equal(actual, before) + + +def _argmax_and_compare( + logits, + token_rows, + proposal_ids, + token_ids_ptrs, + extension_q_offsets, + candidate_active, + entry_sequence_length, + accept_len, + proposal_index, + vocab_size, +): + expected_rows, expected_proposals = draft_argmax_and_store_token_reference( + logits, + token_rows, + extension_q_offsets, + candidate_active, + entry_sequence_length, + accept_len, + proposal_index, + vocab_size, + ) + inputs = [ + logits, + token_ids_ptrs, + extension_q_offsets, + candidate_active, + entry_sequence_length, + accept_len, + ] + inputs_before = [value.clone() for value in inputs] + + result = draft_argmax_and_store_token( + logits, + proposal_ids, + token_ids_ptrs, + extension_q_offsets, + candidate_active, + entry_sequence_length, + accept_len, + proposal_index, + vocab_size, + ) + + assert result is None + assert torch.equal(proposal_ids, expected_proposals) + for actual, expected in zip(token_rows, expected_rows): + if actual is not None: + assert torch.equal(actual, expected) + for actual, before in zip(inputs, inputs_before): + if actual.is_floating_point(): + torch.testing.assert_close(actual, + before, + rtol=0, + atol=0, + equal_nan=True) + else: + assert torch.equal(actual, before) + + +@pytest.mark.parametrize('draft_count', range(1, 8)) +@pytest.mark.parametrize('metrics_enabled', [False, True]) +def test_initialize_target_verification_mixed_rows_and_poison_pointers( + draft_count, + metrics_enabled, +): + batch_size = 7 + generation_count = 5 + + entry_sequence_length = torch.tensor( + [5, 6, 7, 8, 9, 4, 3], + dtype=torch.int32, + device=_DEVICE, + ) + finished_on_entry = torch.tensor( + [False, False, False, True, False, False, False], + dtype=torch.bool, + device=_DEVICE, + ) + speculative_row = torch.tensor( + [True, False, False, True, False, True, False], + dtype=torch.bool, + device=_DEVICE, + ) + request_to_generation_row_offsets = torch.tensor( + [0, 1, 1, 2, 3, 4, 5, 5], + dtype=torch.int32, + device=_DEVICE, + ) + + # Allocate row 5 before row 0 so request order differs from token-row + # allocation order. All ordinary and inherited-finished pointers are null. + row5_backing = torch.full( + (4 + draft_count + 3, ), + _TOKEN_SENTINEL, + dtype=torch.int32, + device=_DEVICE, + ) + row5 = row5_backing[1:-1] + row5.copy_( + torch.arange( + 500, + 500 + row5.numel(), + dtype=torch.int32, + device=_DEVICE, + )) + row0_backing = torch.full( + (5 + draft_count + 3, ), + _TOKEN_SENTINEL, + dtype=torch.int32, + device=_DEVICE, + ) + row0 = row0_backing[1:-1] + row0.copy_( + torch.arange( + 100, + 100 + row0.numel(), + dtype=torch.int32, + device=_DEVICE, + )) + token_rows = [row0, None, None, None, None, row5, None] + request_token_ids_ptrs = _token_pointer_array(token_rows) + + block_logits_active = torch.ones((draft_count + 1, generation_count), + dtype=torch.bool, + device=_DEVICE) + effective_history = torch.full( + (draft_count + 1, generation_count), + _OUTPUT_SENTINEL, + dtype=torch.int32, + device=_DEVICE, + ) + verification_draft_ids = torch.full( + (draft_count, generation_count), + _OUTPUT_SENTINEL, + dtype=torch.int32, + device=_DEVICE, + ) + + accepted_backing = None + accepted_draft_count = None + if metrics_enabled: + accepted_backing, accepted_draft_count = _guarded_int_vector( + [_OUTPUT_SENTINEL] * batch_size) + + expected = initialize_target_verification_reference( + token_rows, + entry_sequence_length, + finished_on_entry, + speculative_row, + request_to_generation_row_offsets, + draft_count, + metrics_enabled, + ) + + pointer_before = request_token_ids_ptrs.clone() + row0_before = row0.clone() + row5_before = row5.clone() + + initialize_target_verification( + block_logits_active, + effective_history, + verification_draft_ids, + request_token_ids_ptrs, + entry_sequence_length, + finished_on_entry, + speculative_row, + accepted_draft_count, + request_to_generation_row_offsets, + ) + + ( + expected_block, + expected_history, + expected_drafts, + expected_accepted_count, + ) = expected + + assert torch.equal(block_logits_active, expected_block) + assert torch.equal(effective_history, expected_history) + assert torch.equal(verification_draft_ids, expected_drafts) + if metrics_enabled: + assert torch.equal(accepted_draft_count, expected_accepted_count) + + if metrics_enabled: + _assert_vector_guards(accepted_backing, _OUTPUT_SENTINEL) + + assert torch.equal(request_token_ids_ptrs, pointer_before) + assert torch.equal(row0, row0_before) + assert torch.equal(row5, row5_before) + _assert_token_guards([row0_backing, row5_backing]) + + +@pytest.mark.parametrize('extension_index', [0, 2, _SPECULATIVE_K - 2]) +@pytest.mark.parametrize('pattern', + ['all_candidate', 'all_non_candidate', 'mixed']) +@pytest.mark.parametrize('batch_size', [0, 1, 255, 256, 257, 513]) +def test_build_draft_extension_key_offsets_matrix(batch_size, pattern, + extension_index): + q_offsets, entry_sequence_length, accept_len = _make_case( + batch_size, pattern) + expected = build_draft_extension_key_offsets_reference( + q_offsets, + entry_sequence_length, + accept_len, + extension_index, + ) + + q_offsets_before = q_offsets.clone() + entry_sequence_length_before = entry_sequence_length.clone() + accept_len_before = accept_len.clone() + backing, k_offsets = _guarded_output(batch_size) + + result = build_draft_extension_key_offsets( + k_offsets, + q_offsets, + entry_sequence_length, + accept_len, + extension_index, + ) + + assert result is None + assert torch.equal(k_offsets, expected) + assert backing[0].item() == _OUTPUT_SENTINEL + assert backing[-1].item() == _OUTPUT_SENTINEL + assert torch.equal(q_offsets, q_offsets_before) + assert torch.equal(entry_sequence_length, entry_sequence_length_before) + assert torch.equal(accept_len, accept_len_before) + + +@pytest.mark.parametrize( + ('case_name', 'entry_sequence_length', 'accept_len', 'extension_index', + 'expected_key_len'), + [ + ('active_final_prompt', 31, 1, 0, 32), + ('active_fallback', 17, 1, 0, 18), + ('speculative_accept_one', 20, 1, 3, 24), + ('speculative_accept_k', 20, _SPECULATIVE_K, 3, + 20 + _SPECULATIVE_K + 3), + ('inherited_finished', 20, 0, 0, 20), + ], +) +def test_build_draft_extension_key_offsets_semantics( + case_name, + entry_sequence_length, + accept_len, + extension_index, + expected_key_len, +): + del case_name + q_offsets = torch.tensor([0, 1], dtype=torch.int32, device=_DEVICE) + entry_sequence_length = torch.tensor([entry_sequence_length], + dtype=torch.int32, + device=_DEVICE) + accept_len = torch.tensor([accept_len], dtype=torch.int32, device=_DEVICE) + k_offsets = torch.full((2, ), + _OUTPUT_SENTINEL, + dtype=torch.int32, + device=_DEVICE) + + build_draft_extension_key_offsets( + k_offsets, + q_offsets, + entry_sequence_length, + accept_len, + extension_index, + ) + + assert torch.equal( + k_offsets, + torch.tensor([0, expected_key_len], dtype=torch.int32, device=_DEVICE)) + + +def test_zero_width_rows_ignore_poisoned_values(): + q_offsets = torch.tensor([0, 1, 1, 2, 2], + dtype=torch.int32, + device=_DEVICE) + entry_sequence_length = torch.tensor( + [10, _POISON_SEQUENCE_LENGTH, 30, _POISON_SEQUENCE_LENGTH], + dtype=torch.int32, + device=_DEVICE, + ) + accept_len = torch.tensor( + [2, _POISON_ACCEPT_LEN, 7, _POISON_ACCEPT_LEN], + dtype=torch.int32, + device=_DEVICE, + ) + extension_index = 4 + k_offsets = torch.empty(5, dtype=torch.int32, device=_DEVICE) + + build_draft_extension_key_offsets( + k_offsets, + q_offsets, + entry_sequence_length, + accept_len, + extension_index, + ) + + expected = torch.tensor([0, 16, 16, 57, 57], + dtype=torch.int32, + device=_DEVICE) + assert torch.equal(k_offsets, expected) + + +def test_build_draft_extension_key_offsets_nondefault_stream(): + batch_size = 257 + extension_index = 2 + q_offsets, entry_sequence_length, accept_len = _make_case( + batch_size, 'mixed') + expected = build_draft_extension_key_offsets_reference( + q_offsets, + entry_sequence_length, + accept_len, + extension_index, + ) + backing, k_offsets = _guarded_output(batch_size) + + current_stream = torch.cuda.current_stream() + stream = torch.cuda.Stream() + stream.wait_stream(current_stream) + + with torch.cuda.stream(stream): + build_draft_extension_key_offsets( + k_offsets, + q_offsets, + entry_sequence_length, + accept_len, + extension_index, + ) + + current_stream.wait_stream(stream) + assert torch.equal(k_offsets, expected) + assert backing[0].item() == _OUTPUT_SENTINEL + assert backing[-1].item() == _OUTPUT_SENTINEL + + +def test_successive_launches_reuse_output_and_preserve_stream_order(): + batch_size = 513 + q_offsets, entry_sequence_length, accept_len = _make_case( + batch_size, 'mixed') + expected_first = build_draft_extension_key_offsets_reference( + q_offsets, + entry_sequence_length, + accept_len, + extension_index=0, + ) + expected_second = build_draft_extension_key_offsets_reference( + q_offsets, + entry_sequence_length, + accept_len, + extension_index=_SPECULATIVE_K - 2, + ) + backing, k_offsets = _guarded_output(batch_size) + + current_stream = torch.cuda.current_stream() + stream = torch.cuda.Stream() + stream.wait_stream(current_stream) + + with torch.cuda.stream(stream): + build_draft_extension_key_offsets( + k_offsets, + q_offsets, + entry_sequence_length, + accept_len, + extension_index=0, + ) + first_snapshot = k_offsets.clone() + + build_draft_extension_key_offsets( + k_offsets, + q_offsets, + entry_sequence_length, + accept_len, + extension_index=_SPECULATIVE_K - 2, + ) + second_snapshot = k_offsets.clone() + + current_stream.wait_stream(stream) + assert torch.equal(first_snapshot, expected_first) + assert torch.equal(second_snapshot, expected_second) + assert torch.equal(k_offsets, expected_second) + assert backing[0].item() == _OUTPUT_SENTINEL + assert backing[-1].item() == _OUTPUT_SENTINEL + + +def test_ordinary_stop_criteria_length_and_packed_phrases(): + backings, allocations = _make_token_rows([ + [1, 2, 3, 0], + [5, 6, 7, 0], + [8, 9, 0, 0], + [11, 12, 13, 0], + [14, 15, 16, 0], + [0, 0, 0, 0], + ]) + token_rows = [ + allocations[2], allocations[0], allocations[4], allocations[1], + allocations[5], allocations[3] + ] + token_ids_ptrs = _token_pointer_array(token_rows) + sequence_length = torch.tensor([2, 3, 3, 3, 0, 3], + dtype=torch.int32, + device=_DEVICE) + limits = torch.tensor([20, 3, 20, 20, 20, 20], + dtype=torch.int32, + device=_DEVICE) + phrases = [ + [[9]], + [], + [[15, 16]], + [[6, 7]], + [], + [[999]], + ] + stop_words = pack_stop_words(phrases, width=4, device=_DEVICE) + finished_backing, finished = _guarded_bool_vector( + [False, False, False, False, False, True]) + + _stop_and_compare(token_rows, token_ids_ptrs, sequence_length, stop_words, + limits, finished) + + assert finished.tolist() == [True, True, True, True, False, True] + _assert_vector_guards(finished_backing, True) + _assert_token_guards(backings) + + +@pytest.mark.parametrize('span_len', range(1, 9)) +def test_speculative_stop_criteria_no_boundary_preserves_span(span_len): + entry_len = 2 + _, token_rows = _make_token_rows([[1, 2] + list(range(20, 20 + span_len)) + + [0]]) + token_ids_ptrs = _token_pointer_array(token_rows) + entry_sequence_length = torch.tensor([entry_len], + dtype=torch.int32, + device=_DEVICE) + accept_len = torch.tensor([span_len], dtype=torch.int32, device=_DEVICE) + limit = torch.tensor([100], dtype=torch.int32, device=_DEVICE) + finished = torch.zeros(1, dtype=torch.bool, device=_DEVICE) + + _stop_and_compare( + token_rows, + token_ids_ptrs, + entry_sequence_length, + None, + limit, + finished, + accept_len=accept_len, + ) + + assert accept_len.item() == span_len + assert not finished.item() + + +@pytest.mark.parametrize('terminal_position', range(8)) +def test_speculative_length_boundary_at_every_position(terminal_position): + entry_len = 3 + _, token_rows = _make_token_rows([[1, 2, 3] + list(range(30, 38)) + [0]]) + token_ids_ptrs = _token_pointer_array(token_rows) + entry_sequence_length = torch.tensor([entry_len], + dtype=torch.int32, + device=_DEVICE) + accept_len = torch.tensor([8], dtype=torch.int32, device=_DEVICE) + limit = torch.tensor([entry_len + terminal_position + 1], + dtype=torch.int32, + device=_DEVICE) + finished = torch.zeros(1, dtype=torch.bool, device=_DEVICE) + + _stop_and_compare( + token_rows, + token_ids_ptrs, + entry_sequence_length, + None, + limit, + finished, + accept_len=accept_len, + ) + + assert accept_len.item() == terminal_position + 1 + assert finished.item() + + +def test_speculative_single_token_stop_at_every_position(): + entry_len = 2 + row_values = [[1, 2] + [100 * b + j for j in range(8)] + [0] + for b in range(8)] + _, token_rows = _make_token_rows(row_values) + token_ids_ptrs = _token_pointer_array(token_rows) + phrases = [[[7000 + b], [100 * b + b], [8000 + b]] for b in range(8)] + stop_words = pack_stop_words(phrases, width=4, device=_DEVICE) + entry_sequence_length = torch.full((8, ), + entry_len, + dtype=torch.int32, + device=_DEVICE) + accept_len = torch.full((8, ), 8, dtype=torch.int32, device=_DEVICE) + limits = torch.full((8, ), 100, dtype=torch.int32, device=_DEVICE) + finished = torch.zeros(8, dtype=torch.bool, device=_DEVICE) + + _stop_and_compare( + token_rows, + token_ids_ptrs, + entry_sequence_length, + stop_words, + limits, + finished, + accept_len=accept_len, + ) + + assert accept_len.tolist() == list(range(1, 9)) + assert finished.all() + + +def test_speculative_stop_phrase_ending_at_every_position_and_crossing_entry(): + entry_len = 2 + row_values = [] + phrases = [] + for b in range(8): + selected = [100 * b + j for j in range(8)] + row = [10 * b + 1, 10 * b + 2] + selected + [0] + phrase = [row[entry_len + b - 1], row[entry_len + b]] + row_values.append(row) + phrases.append([[9999], phrase]) + + _, token_rows = _make_token_rows(row_values) + token_ids_ptrs = _token_pointer_array(token_rows) + stop_words = pack_stop_words(phrases, width=4, device=_DEVICE) + entry_sequence_length = torch.full((8, ), + entry_len, + dtype=torch.int32, + device=_DEVICE) + accept_len = torch.full((8, ), 8, dtype=torch.int32, device=_DEVICE) + limits = torch.full((8, ), 100, dtype=torch.int32, device=_DEVICE) + finished = torch.zeros(8, dtype=torch.bool, device=_DEVICE) + + _stop_and_compare( + token_rows, + token_ids_ptrs, + entry_sequence_length, + stop_words, + limits, + finished, + accept_len=accept_len, + ) + + assert accept_len.tolist() == list(range(1, 9)) + assert finished.all() + + +def test_speculative_stop_edge_semantics_guards_and_effective_eos(): + row_values = [ + [4, 5, 40, 41, 42, 0], + [1, 2, 10, 11, 12, 0], + [1, 2, 20, 21, 22, 0], + [1, 2, 2, 31, 32, 0], + [1, 2, 2, 41, 42, 0], + [1, 2, 50, 51, 52, 53], + [1, 2, 60, 61, 62, 63], + [1, 2, 2, 71, 72, 0], + ] + backings, allocations = _make_token_rows(row_values) + token_rows = [allocations[i] for i in [3, 0, 6, 1, 7, 2, 5, 4]] + token_ids_ptrs = _token_pointer_array(token_rows) + phrases = [ + [], + [[4, 5]], + [], + [], + [[2]], + [[22]], + [[51]], + [], + ] + packed = pack_stop_words(phrases, width=4, device=_DEVICE) + stop_backing = torch.full((packed.numel() + 2, ), + _OUTPUT_SENTINEL, + dtype=torch.int32, + device=_DEVICE) + stop_words = stop_backing[1:-1].view_as(packed) + stop_words.copy_(packed) + entry_sequence_length = torch.full((8, ), + 2, + dtype=torch.int32, + device=_DEVICE) + accept_backing, accept_len = _guarded_int_vector([3, 3, 0, 3, 3, 4, 4, 3]) + limit = torch.tensor([100, 100, 100, 100, 100, 4, 7, 3], + dtype=torch.int32, + device=_DEVICE) + finished_backing, finished = _guarded_bool_vector( + [False, False, False, False, False, False, False, False]) + + _stop_and_compare( + token_rows, + token_ids_ptrs, + entry_sequence_length, + stop_words, + limit, + finished, + accept_len=accept_len, + ) + + assert accept_len.tolist() == [3, 3, 0, 3, 1, 2, 2, 1] + assert finished.tolist() == [ + False, False, False, False, True, True, True, True + ] + assert stop_backing[0].item() == _OUTPUT_SENTINEL + assert stop_backing[-1].item() == _OUTPUT_SENTINEL + _assert_vector_guards(accept_backing, _OUTPUT_SENTINEL) + _assert_vector_guards(finished_backing, True) + _assert_token_guards(backings) + + finished[1] = True + accept_len[1] = 3 + _stop_and_compare( + token_rows, + token_ids_ptrs, + entry_sequence_length, + stop_words, + limit, + finished, + accept_len=accept_len, + ) + assert accept_len[1].item() == 0 + + +def test_stop_criteria_zero_batch(): + token_ids_ptrs = torch.empty(0, dtype=torch.int64, device=_DEVICE) + lengths = torch.empty(0, dtype=torch.int32, device=_DEVICE) + limits = torch.empty(0, dtype=torch.int32, device=_DEVICE) + finished = torch.empty(0, dtype=torch.bool, device=_DEVICE) + accept_len = torch.empty(0, dtype=torch.int32, device=_DEVICE) + + stop_criteria(token_ids_ptrs, lengths, None, limits, finished) + speculative_stop_criteria(token_ids_ptrs, lengths, accept_len, None, + limits, finished) + + +def test_speculative_stop_nondefault_stream_and_successive_snapshots(): + entry_len = 2 + _, token_rows = _make_token_rows([[1, 2, 90, 91, 0]]) + token_ids_ptrs = _token_pointer_array(token_rows) + entry_sequence_length = torch.tensor([entry_len], + dtype=torch.int32, + device=_DEVICE) + accept_len = torch.tensor([2], dtype=torch.int32, device=_DEVICE) + limits = torch.tensor([100], dtype=torch.int32, device=_DEVICE) + finished = torch.zeros(1, dtype=torch.bool, device=_DEVICE) + current_stream = torch.cuda.current_stream() + stream = torch.cuda.Stream() + stream.wait_stream(current_stream) + + with torch.cuda.stream(stream): + speculative_stop_criteria( + token_ids_ptrs, + entry_sequence_length, + accept_len, + None, + limits, + finished, + ) + first_snapshot = accept_len.clone(), finished.clone() + limits.fill_(entry_len + 1) + speculative_stop_criteria( + token_ids_ptrs, + entry_sequence_length, + accept_len, + None, + limits, + finished, + ) + second_snapshot = accept_len.clone(), finished.clone() + + current_stream.wait_stream(stream) + assert first_snapshot[0].item() == 2 + assert not first_snapshot[1].item() + assert second_snapshot[0].item() == 1 + assert second_snapshot[1].item() + + +def test_build_draft_refresh_inputs_mixed_row_modes_mapping_and_guards(): + backings, allocations = _make_token_rows([ + list(range(100, 120)), + list(range(200, 220)), + list(range(300, 320)), + list(range(400, 420)), + list(range(500, 520)), + ]) + token_rows = [ + allocations[2], + allocations[0], + allocations[4], + allocations[1], + allocations[3], + None, + ] + token_ids_ptrs = _token_pointer_array(token_rows) + refresh_q_offsets = torch.tensor( + [0, 3, 5, 7, 11, 12, 12], + dtype=torch.int32, + device=_DEVICE, + ) + refresh_k_offsets = torch.tensor( + [0, 3, 9, 14, 22, 29, 29], + dtype=torch.int32, + device=_DEVICE, + ) + extension_q_offsets = torch.tensor( + [0, 0, 0, 1, 2, 2, 2], + dtype=torch.int32, + device=_DEVICE, + ) + accept_len = torch.tensor( + [ + _POISON_ACCEPT_LEN, + _POISON_ACCEPT_LEN, + _POISON_ACCEPT_LEN, + 2, + 1, + _POISON_ACCEPT_LEN, + ], + dtype=torch.int32, + device=_DEVICE, + ) + limiting = torch.tensor( + [False, False, False, True, True, True], + dtype=torch.bool, + device=_DEVICE, + ) + finished = torch.tensor( + [False, False, False, False, False, True], + dtype=torch.bool, + device=_DEVICE, + ) + draft_backing, draft_input_ids = _guarded_int_vector([_OUTPUT_SENTINEL] * + 12) + selected_backing, selected_token_pos = _guarded_int_vector( + [_OUTPUT_SENTINEL] * 2) + active_backing, candidate_active = _guarded_bool_vector([False, False]) + + _refresh_and_compare( + token_rows, + token_ids_ptrs, + draft_input_ids, + selected_token_pos, + candidate_active, + refresh_q_offsets, + refresh_k_offsets, + extension_q_offsets, + accept_len, + limiting, + finished, + ) + + assert draft_input_ids.tolist() == [ + 301, + 302, + 303, + 105, + 106, + 504, + 505, + 205, + 206, + 0, + 0, + 407, + ] + assert selected_token_pos.tolist() == [6, 8] + assert candidate_active.tolist() == [True, True] + _assert_vector_guards(draft_backing, _OUTPUT_SENTINEL) + _assert_vector_guards(selected_backing, _OUTPUT_SENTINEL) + _assert_vector_guards(active_backing, True) + _assert_token_guards(backings) + + +@pytest.mark.parametrize( + ('width', 'committed'), + [(width, committed) for width in range(2, 9) + for committed in range(width + 1)], +) +def test_build_draft_refresh_inputs_every_speculative_prefix(width, committed): + entry_len = 3 + values = [10, 11, 12] + list(range(100, 100 + width)) + [999] + backings, token_rows = _make_token_rows([values]) + token_ids_ptrs = _token_pointer_array(token_rows) + refresh_q_offsets = torch.tensor([0, width], + dtype=torch.int32, + device=_DEVICE) + refresh_k_offsets = torch.tensor( + [0, entry_len + width - 1], + dtype=torch.int32, + device=_DEVICE, + ) + extension_q_offsets = torch.tensor([0, 1], + dtype=torch.int32, + device=_DEVICE) + accept_len = torch.tensor([committed], dtype=torch.int32, device=_DEVICE) + limiting = torch.ones(1, dtype=torch.bool, device=_DEVICE) + finished = torch.zeros(1, dtype=torch.bool, device=_DEVICE) + draft_input_ids = torch.full( + (width, ), + _OUTPUT_SENTINEL, + dtype=torch.int32, + device=_DEVICE, + ) + selected_token_pos = torch.full( + (1, ), + _OUTPUT_SENTINEL, + dtype=torch.int32, + device=_DEVICE, + ) + candidate_active = torch.ones(1, dtype=torch.bool, device=_DEVICE) + + _refresh_and_compare( + token_rows, + token_ids_ptrs, + draft_input_ids, + selected_token_pos, + candidate_active, + refresh_q_offsets, + refresh_k_offsets, + extension_q_offsets, + accept_len, + limiting, + finished, + ) + + assert draft_input_ids[:committed].tolist() == list( + range(100, 100 + committed)) + assert not draft_input_ids[committed:].any() + assert candidate_active.item() == (committed > 0) + assert selected_token_pos.item() == max(0, committed - 1) + _assert_token_guards(backings) + + +def test_build_draft_refresh_inputs_terminal_and_inherited_finished_rows(): + backings, allocations = _make_token_rows([ + [1, 2, 3, 40, 41, 42, 43, 44], + ]) + token_rows = [allocations[0], None] + token_ids_ptrs = _token_pointer_array(token_rows) + refresh_q_offsets = torch.tensor([0, 4, 8], + dtype=torch.int32, + device=_DEVICE) + refresh_k_offsets = torch.tensor([0, 6, 12], + dtype=torch.int32, + device=_DEVICE) + extension_q_offsets = torch.tensor([0, 1, 2], + dtype=torch.int32, + device=_DEVICE) + accept_len = torch.tensor([3, 0], dtype=torch.int32, device=_DEVICE) + limiting = torch.ones(2, dtype=torch.bool, device=_DEVICE) + finished = torch.ones(2, dtype=torch.bool, device=_DEVICE) + draft_input_ids = torch.full( + (8, ), + _OUTPUT_SENTINEL, + dtype=torch.int32, + device=_DEVICE, + ) + selected_token_pos = torch.full( + (2, ), + _OUTPUT_SENTINEL, + dtype=torch.int32, + device=_DEVICE, + ) + candidate_active = torch.ones(2, dtype=torch.bool, device=_DEVICE) + + _refresh_and_compare( + token_rows, + token_ids_ptrs, + draft_input_ids, + selected_token_pos, + candidate_active, + refresh_q_offsets, + refresh_k_offsets, + extension_q_offsets, + accept_len, + limiting, + finished, + ) + + assert draft_input_ids.tolist() == [40, 41, 42, 0, 0, 0, 0, 0] + assert selected_token_pos.tolist() == [0, 0] + assert candidate_active.tolist() == [False, False] + _assert_token_guards(backings) + + +@pytest.mark.parametrize('committed', [0, 1]) +def test_build_draft_refresh_inputs_ordinary_fallback(committed): + if committed: + backings, token_rows = _make_token_rows([[1, 2, 3, 4, 5, 91]]) + else: + backings, token_rows = [], [None] + token_ids_ptrs = _token_pointer_array(token_rows) + draft_backing, draft_input_ids = _guarded_int_vector([_OUTPUT_SENTINEL]) + selected_backing, selected_token_pos = _guarded_int_vector([]) + active_backing, candidate_active = _guarded_bool_vector([]) + + _refresh_and_compare( + token_rows, + token_ids_ptrs, + draft_input_ids, + selected_token_pos, + candidate_active, + torch.tensor([0, 1], dtype=torch.int32, device=_DEVICE), + torch.tensor([0, 5], dtype=torch.int32, device=_DEVICE), + torch.tensor([0, 0], dtype=torch.int32, device=_DEVICE), + torch.tensor([committed], dtype=torch.int32, device=_DEVICE), + torch.ones(1, dtype=torch.bool, device=_DEVICE), + torch.tensor([not committed], dtype=torch.bool, device=_DEVICE), + ) + + assert draft_input_ids.item() == (91 if committed else 0) + _assert_vector_guards(draft_backing, _OUTPUT_SENTINEL) + _assert_vector_guards(selected_backing, _OUTPUT_SENTINEL) + _assert_vector_guards(active_backing, True) + _assert_token_guards(backings) + + +def test_build_draft_refresh_inputs_empty_batches_and_zero_packed_rows(): + empty_i32 = torch.empty(0, dtype=torch.int32, device=_DEVICE) + empty_i64 = torch.empty(0, dtype=torch.int64, device=_DEVICE) + empty_bool = torch.empty(0, dtype=torch.bool, device=_DEVICE) + zero_offsets = torch.zeros(1, dtype=torch.int32, device=_DEVICE) + + build_draft_refresh_inputs( + empty_i32, + empty_i32, + empty_bool, + empty_i64, + zero_offsets, + zero_offsets, + zero_offsets, + empty_i32, + empty_bool, + empty_bool, + ) + + batch_size = 4 + pointer_array = torch.zeros(batch_size, dtype=torch.int64, device=_DEVICE) + offsets = torch.zeros(batch_size + 1, dtype=torch.int32, device=_DEVICE) + accept_len = torch.full( + (batch_size, ), + _POISON_ACCEPT_LEN, + dtype=torch.int32, + device=_DEVICE, + ) + limiting = torch.ones(batch_size, dtype=torch.bool, device=_DEVICE) + finished = torch.ones(batch_size, dtype=torch.bool, device=_DEVICE) + build_draft_refresh_inputs( + empty_i32, + empty_i32, + empty_bool, + pointer_array, + offsets, + torch.full_like(offsets, _POISON_SEQUENCE_LENGTH), + offsets, + accept_len, + limiting, + finished, + ) + + +def test_build_draft_refresh_inputs_nondefault_stream_successive_snapshots(): + _, token_rows = _make_token_rows([[1, 2, 3, 31, 32, 33, 34]]) + token_ids_ptrs = _token_pointer_array(token_rows) + q_offsets = torch.tensor([0, 3], dtype=torch.int32, device=_DEVICE) + k_offsets = torch.tensor([0, 5], dtype=torch.int32, device=_DEVICE) + extension_offsets = torch.tensor([0, 1], dtype=torch.int32, device=_DEVICE) + accept_len = torch.tensor([1], dtype=torch.int32, device=_DEVICE) + limiting = torch.ones(1, dtype=torch.bool, device=_DEVICE) + finished = torch.zeros(1, dtype=torch.bool, device=_DEVICE) + draft_input_ids = torch.empty(3, dtype=torch.int32, device=_DEVICE) + selected_token_pos = torch.empty(1, dtype=torch.int32, device=_DEVICE) + candidate_active = torch.empty(1, dtype=torch.bool, device=_DEVICE) + current_stream = torch.cuda.current_stream() + stream = torch.cuda.Stream() + stream.wait_stream(current_stream) + + with torch.cuda.stream(stream): + build_draft_refresh_inputs( + draft_input_ids, + selected_token_pos, + candidate_active, + token_ids_ptrs, + q_offsets, + k_offsets, + extension_offsets, + accept_len, + limiting, + finished, + ) + first_snapshot = ( + draft_input_ids.clone(), + selected_token_pos.clone(), + candidate_active.clone(), + ) + accept_len.fill_(3) + build_draft_refresh_inputs( + draft_input_ids, + selected_token_pos, + candidate_active, + token_ids_ptrs, + q_offsets, + k_offsets, + extension_offsets, + accept_len, + limiting, + finished, + ) + second_snapshot = ( + draft_input_ids.clone(), + selected_token_pos.clone(), + candidate_active.clone(), + ) + + current_stream.wait_stream(stream) + assert first_snapshot[0].tolist() == [31, 0, 0] + assert first_snapshot[1].item() == 0 + assert first_snapshot[2].item() + assert second_snapshot[0].tolist() == [31, 32, 33] + assert second_snapshot[1].item() == 2 + assert second_snapshot[2].item() + + +@pytest.mark.parametrize('dtype', [torch.float16, torch.bfloat16]) +@pytest.mark.parametrize('proposal_index', range(_SPECULATIVE_K)) +def test_draft_argmax_random_rows_every_proposal_index(dtype, proposal_index): + backings, allocations = _make_token_rows([ + list(range(100, 132)), + list(range(200, 232)), + list(range(300, 332)), + ]) + token_rows = [allocations[2], None, allocations[0], None, allocations[1]] + token_ids_ptrs = _token_pointer_array(token_rows) + extension_offsets = torch.tensor( + [0, 1, 1, 2, 2, 3], + dtype=torch.int32, + device=_DEVICE, + ) + candidate_active = torch.ones(3, dtype=torch.bool, device=_DEVICE) + entry_lengths = torch.tensor( + [4, _POISON_SEQUENCE_LENGTH, 7, _POISON_SEQUENCE_LENGTH, 9], + dtype=torch.int32, + device=_DEVICE, + ) + accept_len = torch.tensor( + [2, _POISON_ACCEPT_LEN, 3, _POISON_ACCEPT_LEN, 1], + dtype=torch.int32, + device=_DEVICE, + ) + vocab_size = 37 + logits = torch.randn(3, vocab_size + 5, dtype=dtype, device=_DEVICE) + logits[:, vocab_size:] = torch.inf + proposal_backing, proposal_ids = _guarded_int_vector([_OUTPUT_SENTINEL] * + 3) + + _argmax_and_compare( + logits, + token_rows, + proposal_ids, + token_ids_ptrs, + extension_offsets, + candidate_active, + entry_lengths, + accept_len, + proposal_index, + vocab_size, + ) + + _assert_vector_guards(proposal_backing, _OUTPUT_SENTINEL) + _assert_token_guards(backings) + + +@pytest.mark.parametrize('dtype', [torch.float16, torch.bfloat16]) +def test_draft_argmax_concrete_vocab_size(dtype): + vocab_size = 151936 + backings, token_rows = _make_token_rows([list(range(32))]) + token_ids_ptrs = _token_pointer_array(token_rows) + logits = torch.randn( + 1, + vocab_size + 8, + dtype=dtype, + device=_DEVICE, + ) + logits[:, vocab_size:] = torch.inf + proposal_ids = torch.full( + (1, ), + _OUTPUT_SENTINEL, + dtype=torch.int32, + device=_DEVICE, + ) + + _argmax_and_compare( + logits, + token_rows, + proposal_ids, + token_ids_ptrs, + torch.tensor([0, 1], dtype=torch.int32, device=_DEVICE), + torch.ones(1, dtype=torch.bool, device=_DEVICE), + torch.tensor([5], dtype=torch.int32, device=_DEVICE), + torch.tensor([3], dtype=torch.int32, device=_DEVICE), + 2, + vocab_size, + ) + + _assert_token_guards(backings) + + +@pytest.mark.parametrize('dtype', [torch.float16, torch.bfloat16]) +def test_draft_argmax_ties_infinities_signed_zero_and_nan(dtype): + values = [ + [-9.0, 4.0, 1.0, 4.0, -2.0, -3.0], + [0.0, torch.inf, 5.0, torch.inf, -1.0, -2.0], + [-3.0, -0.0, 0.0, -2.0, -4.0, -5.0], + [-torch.inf] * 6, + [torch.nan] * 6, + [torch.nan, -2.0, 7.0, torch.nan, 7.0, -torch.inf], + ] + expected_token_ids = [1, 1, 1, 0, 0, 2] + logits = torch.tensor(values, dtype=dtype, device=_DEVICE) + logits = torch.cat([ + logits, + torch.full( + (len(values), 3), + torch.inf, + dtype=dtype, + device=_DEVICE, + ), + ], + dim=1) + backings, token_rows = _make_token_rows([list(range(20)) for _ in values]) + token_ids_ptrs = _token_pointer_array(token_rows) + batch_size = len(values) + offsets = torch.arange( + batch_size + 1, + dtype=torch.int32, + device=_DEVICE, + ) + proposal_backing, proposal_ids = _guarded_int_vector([_OUTPUT_SENTINEL] * + batch_size) + + _argmax_and_compare( + logits, + token_rows, + proposal_ids, + token_ids_ptrs, + offsets, + torch.ones(batch_size, dtype=torch.bool, device=_DEVICE), + torch.full((batch_size, ), 4, dtype=torch.int32, device=_DEVICE), + torch.arange(batch_size, dtype=torch.int32, device=_DEVICE), + 1, + 6, + ) + + assert proposal_ids.tolist() == expected_token_ids + _assert_vector_guards(proposal_backing, _OUTPUT_SENTINEL) + _assert_token_guards(backings) + + +def test_draft_argmax_inactive_and_noncandidate_poison_rows(): + backings, allocations = _make_token_rows([ + list(range(40)), + list(range(100, 140)), + ]) + token_rows = [allocations[1], None, None, None, None, allocations[0]] + token_ids_ptrs = _token_pointer_array(token_rows) + extension_offsets = torch.tensor( + [0, 1, 1, 2, 2, 2, 3], + dtype=torch.int32, + device=_DEVICE, + ) + candidate_active = torch.tensor( + [True, False, True], + dtype=torch.bool, + device=_DEVICE, + ) + entry_lengths = torch.tensor( + [ + 5, + _POISON_SEQUENCE_LENGTH, + _POISON_SEQUENCE_LENGTH, + _POISON_SEQUENCE_LENGTH, + _POISON_SEQUENCE_LENGTH, + 8, + ], + dtype=torch.int32, + device=_DEVICE, + ) + accept_len = torch.tensor( + [ + 2, + _POISON_ACCEPT_LEN, + _POISON_ACCEPT_LEN, + _POISON_ACCEPT_LEN, + _POISON_ACCEPT_LEN, + 3, + ], + dtype=torch.int32, + device=_DEVICE, + ) + logits = torch.tensor( + [ + [1.0, 9.0, 3.0, 2.0], + [torch.nan, torch.nan, torch.nan, torch.nan], + [8.0, 2.0, 11.0, 5.0], + ], + dtype=torch.float16, + device=_DEVICE, + ) + proposal_backing, proposal_ids = _guarded_int_vector([_OUTPUT_SENTINEL] * + 3) + + _argmax_and_compare( + logits, + token_rows, + proposal_ids, + token_ids_ptrs, + extension_offsets, + candidate_active, + entry_lengths, + accept_len, + 2, + 4, + ) + + assert proposal_ids.tolist() == [1, 0, 2] + _assert_vector_guards(proposal_backing, _OUTPUT_SENTINEL) + _assert_token_guards(backings) + + +def test_draft_argmax_empty_batch_and_zero_candidates(): + empty_i32 = torch.empty(0, dtype=torch.int32, device=_DEVICE) + empty_i64 = torch.empty(0, dtype=torch.int64, device=_DEVICE) + empty_bool = torch.empty(0, dtype=torch.bool, device=_DEVICE) + empty_logits = torch.empty( + (0, 8), + dtype=torch.float16, + device=_DEVICE, + ) + zero_offsets = torch.zeros(1, dtype=torch.int32, device=_DEVICE) + + draft_argmax_and_store_token( + empty_logits, + empty_i32, + empty_i64, + zero_offsets, + empty_bool, + empty_i32, + empty_i32, + 0, + 8, + ) + + batch_size = 4 + draft_argmax_and_store_token( + empty_logits, + empty_i32, + torch.zeros(batch_size, dtype=torch.int64, device=_DEVICE), + torch.zeros(batch_size + 1, dtype=torch.int32, device=_DEVICE), + empty_bool, + torch.full( + (batch_size, ), + _POISON_SEQUENCE_LENGTH, + dtype=torch.int32, + device=_DEVICE, + ), + torch.full( + (batch_size, ), + _POISON_ACCEPT_LEN, + dtype=torch.int32, + device=_DEVICE, + ), + 0, + 8, + ) + + +def test_draft_argmax_nondefault_stream_successive_proposal_indices(): + _, token_rows = _make_token_rows([[10] * 24]) + token_ids_ptrs = _token_pointer_array(token_rows) + offsets = torch.tensor([0, 1], dtype=torch.int32, device=_DEVICE) + active = torch.ones(1, dtype=torch.bool, device=_DEVICE) + entry_lengths = torch.tensor([4], dtype=torch.int32, device=_DEVICE) + accept_len = torch.tensor([2], dtype=torch.int32, device=_DEVICE) + logits = torch.tensor( + [[1.0, 2.0, 9.0, 3.0]], + dtype=torch.bfloat16, + device=_DEVICE, + ) + proposal_ids = torch.empty(1, dtype=torch.int32, device=_DEVICE) + current_stream = torch.cuda.current_stream() + stream = torch.cuda.Stream() + stream.wait_stream(current_stream) + + with torch.cuda.stream(stream): + draft_argmax_and_store_token( + logits, + proposal_ids, + token_ids_ptrs, + offsets, + active, + entry_lengths, + accept_len, + 0, + 4, + ) + first_snapshot = proposal_ids.clone(), token_rows[0].clone() + logits.copy_( + torch.tensor( + [[8.0, 2.0, 1.0, 11.0]], + dtype=torch.bfloat16, + device=_DEVICE, + )) + draft_argmax_and_store_token( + logits, + proposal_ids, + token_ids_ptrs, + offsets, + active, + entry_lengths, + accept_len, + 1, + 4, + ) + second_snapshot = proposal_ids.clone(), token_rows[0].clone() + + current_stream.wait_stream(stream) + assert first_snapshot[0].item() == 2 + assert first_snapshot[1][6].item() == 2 + assert second_snapshot[0].item() == 3 + assert second_snapshot[1][6].item() == 2 + assert second_snapshot[1][7].item() == 3 diff --git a/tests/turbomind/target_hidden_projection/__init__.py b/tests/turbomind/target_hidden_projection/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/tests/turbomind/target_hidden_projection/reference.py b/tests/turbomind/target_hidden_projection/reference.py new file mode 100644 index 0000000000..7e7d04cd77 --- /dev/null +++ b/tests/turbomind/target_hidden_projection/reference.py @@ -0,0 +1,17 @@ +from __future__ import annotations + +import torch + + +def capture_target_hidden_rows( + packed_residual: torch.Tensor, + captured: torch.Tensor, + owned_begin: int, + owned_row_count: int, + tap_ordinal: int, +) -> None: + """Apply the target-row capture contract using direct Torch assignment.""" + hidden_units = packed_residual.shape[1] + tap_begin = tap_ordinal * hidden_units + captured[:owned_row_count, tap_begin:tap_begin + hidden_units] = ( + packed_residual[owned_begin:owned_begin + owned_row_count, :]) diff --git a/tests/turbomind/target_hidden_projection/target_hidden_projection.py b/tests/turbomind/target_hidden_projection/target_hidden_projection.py new file mode 100644 index 0000000000..e073eb3be8 --- /dev/null +++ b/tests/turbomind/target_hidden_projection/target_hidden_projection.py @@ -0,0 +1,46 @@ +from __future__ import annotations + +import torch + +_NATIVE_SYMBOL = 'capture_target_hidden_rows' + + +def _load_native_bridge(): + try: + import _turbomind as tm + except ImportError: + return None + return tm if hasattr(tm, _NATIVE_SYMBOL) else None + + +def _require_native_bridge(): + tm = _load_native_bridge() + if tm is None: + raise ImportError( + 'TurboMind target-hidden capture bridge is unavailable; ' + f'required symbol: {_NATIVE_SYMBOL}') + return tm + + +def is_available() -> bool: + return _load_native_bridge() is not None + + +def capture_target_hidden_rows( + packed_residual: torch.Tensor, + captured: torch.Tensor, + owned_begin: int, + owned_row_count: int, + tap_ordinal: int, +) -> None: + stream_ptr = int( + torch.cuda.current_stream( + packed_residual.device).cuda_stream) + _require_native_bridge().capture_target_hidden_rows( + packed_residual, + captured, + int(owned_begin), + int(owned_row_count), + int(tap_ordinal), + stream_ptr, + ) diff --git a/tests/turbomind/target_hidden_projection/test_target_hidden_projection.py b/tests/turbomind/target_hidden_projection/test_target_hidden_projection.py new file mode 100644 index 0000000000..f3d97d1cd0 --- /dev/null +++ b/tests/turbomind/target_hidden_projection/test_target_hidden_projection.py @@ -0,0 +1,330 @@ +from __future__ import annotations + +from dataclasses import dataclass + +import pytest +import torch + +from .reference import capture_target_hidden_rows as reference_capture +from .target_hidden_projection import ( + capture_target_hidden_rows as native_capture, +) +from .target_hidden_projection import is_available + +pytestmark = pytest.mark.skipif( + not torch.cuda.is_available(), + reason='CUDA is required for target-hidden capture tests', +) + +_CAPTURE_POISON_BITS = 0x3555 +_SOURCE_POISON_BITS = 0x5A5A + + +@dataclass(frozen=True) +class _PitchedTensor: + storage: torch.Tensor + view: torch.Tensor + + +def _require_bridge() -> None: + if not is_available(): + pytest.skip('TurboMind target-hidden capture bridge is unavailable') + + +def _pitched_tensor( + rows: int, + width: int, + *, + row_padding: int, + prefix: int, + suffix: int, + dtype: torch.dtype, + poison_bits: int, +) -> _PitchedTensor: + leading_dimension = width + row_padding + storage = torch.empty( + prefix + rows * leading_dimension + suffix, + dtype=dtype, + device='cuda', + ) + storage.view(torch.int16).fill_(poison_bits) + view = torch.as_strided( + storage, + (rows, width), + (leading_dimension, 1), + storage_offset=prefix, + ) + return _PitchedTensor(storage, view) + + +def _row_bits(hidden_units: int, seed: int) -> torch.Tensor: + values = ( + torch.arange( + hidden_units, + dtype=torch.int32, + device='cuda', + ) * 251 + seed + ) + return values.remainder(65536).sub(32768).to(torch.int16) + + +def _make_source( + rows: int, + hidden_units: int, + owned_begin: int, + owned_row_count: int, + dtype: torch.dtype, + *, + seed: int, +) -> _PitchedTensor: + source = _pitched_tensor( + rows, + hidden_units, + row_padding=7, + prefix=5, + suffix=11, + dtype=dtype, + poison_bits=_SOURCE_POISON_BITS, + ) + for row in range(owned_begin, owned_begin + owned_row_count): + source.view[row].view(torch.int16).copy_( + _row_bits(hidden_units, seed + row * 977)) + return source + + +def _make_capture( + capacity: int, + hidden_units: int, + tap_count: int, + dtype: torch.dtype, +) -> _PitchedTensor: + return _pitched_tensor( + capacity, + tap_count * hidden_units, + row_padding=13, + prefix=9, + suffix=17, + dtype=dtype, + poison_bits=_CAPTURE_POISON_BITS, + ) + + +def _assert_bits_equal(actual: torch.Tensor, expected: torch.Tensor) -> None: + assert actual.dtype == expected.dtype + assert actual.shape == expected.shape + assert torch.equal( + actual.view(torch.int16), + expected.view(torch.int16), + ) + + +def _run_capture_case( + *, + dtype: torch.dtype, + source_rows: int, + owned_begin: int, + owned_row_count: int, + capacity: int, + hidden_units: int, + tap_count: int, + tap_ordinal: int, + seed: int = 101, +) -> None: + _require_bridge() + source = _make_source( + source_rows, + hidden_units, + owned_begin, + owned_row_count, + dtype, + seed=seed, + ) + source_before = source.storage.clone() + captured = _make_capture( + capacity, + hidden_units, + tap_count, + dtype, + ) + expected_storage = captured.storage.clone() + expected = torch.as_strided( + expected_storage, + captured.view.shape, + captured.view.stride(), + storage_offset=captured.view.storage_offset(), + ) + + reference_capture( + source.view, + expected, + owned_begin, + owned_row_count, + tap_ordinal, + ) + native_capture( + source.view, + captured.view, + owned_begin, + owned_row_count, + tap_ordinal, + ) + + _assert_bits_equal(captured.storage, expected_storage) + _assert_bits_equal(source.storage, source_before) + + +@pytest.mark.parametrize('dtype', [torch.float16, torch.bfloat16]) +@pytest.mark.parametrize( + 'source_rows,owned_begin,owned_row_count,capacity,hidden_units,tap_ordinal', + [ + (0, 0, 0, 4, 13, 2), + (6, 3, 1, 3, 17, 0), + (11, 2, 5, 7, 31, 3), + (7, 0, 7, 7, 16, 4), + ], + ids=['zero', 'one', 'uneven', 'full'], +) +def test_capture_interval_matrix_is_bit_exact( + dtype, + source_rows, + owned_begin, + owned_row_count, + capacity, + hidden_units, + tap_ordinal, +): + _run_capture_case( + dtype=dtype, + source_rows=source_rows, + owned_begin=owned_begin, + owned_row_count=owned_row_count, + capacity=capacity, + hidden_units=hidden_units, + tap_count=5, + tap_ordinal=tap_ordinal, + ) + + +@pytest.mark.parametrize('dtype', [torch.float16, torch.bfloat16]) +@pytest.mark.parametrize('tap_ordinal', range(5)) +def test_each_tap_isolated(dtype, tap_ordinal): + _run_capture_case( + dtype=dtype, + source_rows=9, + owned_begin=3, + owned_row_count=4, + capacity=6, + hidden_units=23, + tap_count=5, + tap_ordinal=tap_ordinal, + seed=1000 + tap_ordinal, + ) + + +@pytest.mark.parametrize('dtype', [torch.float16, torch.bfloat16]) +def test_all_five_taps_populated(dtype): + _require_bridge() + hidden_units = 29 + owned_begin = 2 + owned_row_count = 5 + captured = _make_capture(7, hidden_units, 5, dtype) + expected_storage = captured.storage.clone() + expected = torch.as_strided( + expected_storage, + captured.view.shape, + captured.view.stride(), + storage_offset=captured.view.storage_offset(), + ) + + sources = [] + source_snapshots = [] + for tap_ordinal in range(5): + source = _make_source( + 10, + hidden_units, + owned_begin, + owned_row_count, + dtype, + seed=2000 + 100 * tap_ordinal, + ) + sources.append(source) + source_snapshots.append(source.storage.clone()) + reference_capture( + source.view, + expected, + owned_begin, + owned_row_count, + tap_ordinal, + ) + native_capture( + source.view, + captured.view, + owned_begin, + owned_row_count, + tap_ordinal, + ) + + _assert_bits_equal(captured.storage, expected_storage) + for source, snapshot in zip(sources, source_snapshots): + _assert_bits_equal(source.storage, snapshot) + + +@pytest.mark.parametrize('dtype', [torch.float16, torch.bfloat16]) +def test_single_tap_diagnostic(dtype): + _run_capture_case( + dtype=dtype, + source_rows=8, + owned_begin=1, + owned_row_count=6, + capacity=6, + hidden_units=37, + tap_count=1, + tap_ordinal=0, + ) + + +@pytest.mark.parametrize('dtype', [torch.float16, torch.bfloat16]) +def test_hidden_size_4096(dtype): + _run_capture_case( + dtype=dtype, + source_rows=6, + owned_begin=2, + owned_row_count=3, + capacity=4, + hidden_units=4096, + tap_count=5, + tap_ordinal=4, + ) + + +def test_nondefault_stream_and_asynchronous_return(): + _require_bridge() + source = _make_source( + 8, + 41, + 2, + 4, + torch.float16, + seed=31415, + ) + captured = _make_capture(6, 41, 5, torch.float16) + expected_storage = captured.storage.clone() + expected = torch.as_strided( + expected_storage, + captured.view.shape, + captured.view.stride(), + storage_offset=captured.view.storage_offset(), + ) + reference_capture(source.view, expected, 2, 4, 3) + + torch.cuda.synchronize() + stream = torch.cuda.Stream(device=source.view.device) + with torch.cuda.stream(stream): + if hasattr(torch.cuda, '_sleep'): + torch.cuda._sleep(1_000_000_000) + native_capture(source.view, captured.view, 2, 4, 3) + if hasattr(torch.cuda, '_sleep'): + assert not stream.query() + stream.synchronize() + + _assert_bits_equal(captured.storage, expected_storage) From ffb4c96eec862874fcfb9a37a87482bc17018a89 Mon Sep 17 00:00:00 2001 From: bltcn Date: Sat, 3 Oct 2026 10:13:42 +0800 Subject: [PATCH 4/7] fix(turbomind): normalize MTP speculative method aliases to registered 'mtp' The CLI exposes qwen3_5_mtp/hy3_mtp/deepseek_mtp, but TurboMind registers the MTP head under the C++ name "mtp" (TM_REGISTER_SPECULATIVE_MODEL("mtp", ...)). Without normalization, --speculative-algorithm qwen3_5_mtp hit build_draft_model's ValueError and the C++ 'unknown speculative method' check. Add normalize_spec_method() and apply it at both the draft-model builder and the EngineConfig.spec_method assignment so the registered name reaches C++. --- lmdeploy/turbomind/spec_decode.py | 15 ++++++++++++++- lmdeploy/turbomind/turbomind.py | 4 ++-- 2 files changed, 16 insertions(+), 3 deletions(-) diff --git a/lmdeploy/turbomind/spec_decode.py b/lmdeploy/turbomind/spec_decode.py index 243eae610e..da28467951 100644 --- a/lmdeploy/turbomind/spec_decode.py +++ b/lmdeploy/turbomind/spec_decode.py @@ -29,13 +29,26 @@ class DraftWeightSpec: } +# TurboMind registers MTP heads under the C++ name "mtp" (see +# TM_REGISTER_SPECULATIVE_MODEL("mtp", ...)); the CLI exposes model-specific +# aliases. Normalize so --speculative-algorithm qwen3_5_mtp (and the other MTP +# aliases) reaches the registered MTP model. Must stay in sync with the +# registered speculative-model names in src/turbomind/models/speculative/. +_MTP_ALIASES = {'qwen3_5_mtp': 'mtp', 'hy3_mtp': 'mtp', 'deepseek_mtp': 'mtp'} + + +def normalize_spec_method(method: str) -> str: + """Map a CLI speculative-algorithm name to the TurboMind-registered name.""" + return _MTP_ALIASES.get(method, method) + + def build_draft_model(speculative_config, target_model, target_model_path, engine_data_type, download_dir=None): """Build the draft weight mapper and resolve the checkpoint it reads.""" - method = speculative_config.method + method = normalize_spec_method(speculative_config.method) if method == 'mtp': target_text = target_model.text_model diff --git a/lmdeploy/turbomind/turbomind.py b/lmdeploy/turbomind/turbomind.py index acf3b539cb..3abd6eda4c 100644 --- a/lmdeploy/turbomind/turbomind.py +++ b/lmdeploy/turbomind/turbomind.py @@ -237,7 +237,7 @@ def _from_hf(self, from .converter import get_tm_config from .model_loader import ModelLoader - from .spec_decode import build_draft_model + from .spec_decode import build_draft_model, normalize_spec_method model, model_path, data_type = get_tm_config(model_path, engine_config, trust_remote_code=trust_remote_code) @@ -253,7 +253,7 @@ def _from_hf(self, draft_model, draft_model_path = build_draft_model( speculative_config, model, model_path, data_type, engine_config.download_dir) - spec_method = speculative_config.method + spec_method = normalize_spec_method(speculative_config.method) spec_num_draft_tokens = speculative_config.num_speculative_tokens spec_tap_layer_ids = list(draft_model.tap_layer_ids) From 00feb403c4ce11d4d21240cb6d9e57204b2af1c5 Mon Sep 17 00:00:00 2001 From: bltcn Date: Sat, 3 Oct 2026 15:30:08 +0800 Subject: [PATCH 5/7] fix(turbomind): GDN state commit kernel missing half_t on pre-SM80 (MTP on 2080Ti) ForwardSpeculativeRound -> CommitAcceptedRecurrentStateKernel dispatches the recurrent-state dtype through TM_DISPATCH_DTYPES with only (bfloat16_t, float). On pre-SM80 GPUs (no bf16 tensor core) the engine dtype is float16, so the GDN recurrent state is f16 and the speculative round aborts with 'unsupported type: f16'. The kernel body computes in float and only touches StateT via ToFloat/FromFloat, and the input side already dispatches half_t, so adding half_t is a correctness fix for SM75 MTP. Not caught by CI (A100/bf16 path never hits this branch). --- src/turbomind/kernels/linear_attn/gdn_state_transaction.cu | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/src/turbomind/kernels/linear_attn/gdn_state_transaction.cu b/src/turbomind/kernels/linear_attn/gdn_state_transaction.cu index 9fbe6dfec8..d38b5bd253 100644 --- a/src/turbomind/kernels/linear_attn/gdn_state_transaction.cu +++ b/src/turbomind/kernels/linear_attn/gdn_state_transaction.cu @@ -382,7 +382,11 @@ void invokeCommitAcceptedRecurrentState( args.heads_per_block, total_work); }; - TM_DISPATCH_DTYPES(state_dtype, launch_state, bfloat16_t, float); + // f16 is required on pre-SM80 (Turing/Maxwell, no bf16 tensor core), where + // the recurrent state follows the engine's float16 dtype. The kernel body is + // computed in float and only touches StateT via ToFloat/FromFloat, so half_t + // is safe here (mirrors the input-side dispatch below). + TM_DISPATCH_DTYPES(state_dtype, launch_state, half_t, bfloat16_t, float); }; TM_DISPATCH_DTYPES(args.key.dtype(), launch_input, half_t, bfloat16_t); TM_CUDA_CHECK(cudaGetLastError()); From f0a42dff391ea37972cbe1abaaa5027d8c5939c6 Mon Sep 17 00:00:00 2001 From: bltcn Date: Sat, 3 Oct 2026 20:34:27 +0800 Subject: [PATCH 6/7] fix(turbomind): keep normalize_spec_method docstring within summary width CI lint's docformatter hook (v1.7.7, --wrap-descriptions 120) rejects the 82-char one-line summary. Shortened to 58 chars; ruff + docformatter clean. --- lmdeploy/turbomind/spec_decode.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/lmdeploy/turbomind/spec_decode.py b/lmdeploy/turbomind/spec_decode.py index da28467951..708aa715e6 100644 --- a/lmdeploy/turbomind/spec_decode.py +++ b/lmdeploy/turbomind/spec_decode.py @@ -38,7 +38,7 @@ class DraftWeightSpec: def normalize_spec_method(method: str) -> str: - """Map a CLI speculative-algorithm name to the TurboMind-registered name.""" + """Normalize a CLI speculative name to the TurboMind name.""" return _MTP_ALIASES.get(method, method) From ddfab39d3a5e78bfe99e36783ed88c5b76ae3d12 Mon Sep 17 00:00:00 2001 From: bltcn Date: Sat, 3 Oct 2026 21:24:49 +0800 Subject: [PATCH 7/7] fix(tests): use canonical package import for _turbomind + ruff lint fixes Inherited CI blockers from #5006 (this PR is stacked on it): - unit_test (exit 2, collection error): 3 top-level 'import _turbomind' fail under CI 'pip install -e .' because the bare module name is not on sys.path. Switch to 'from lmdeploy.turbomind import _tm' (canonical, matches test_linear.py) and function-level imports to the package name. - lint (ruff): F841 dead var, F401 unused imports, I001 import sorting. Verified in CI-matched env (0.18.0 + cu12.8 + built _tm.so + torch): 308 tests collected, 0 collection errors (was 213 collected, 2 errors). --- .../attention/test_verification_attention.py | 9 +-------- tests/turbomind/attention/verification_attention.py | 3 ++- tests/turbomind/draft_carry/draft_carry.py | 2 +- tests/turbomind/draft_carry/test_draft_carry.py | 8 ++++---- .../speculative_sampling/speculative_sampling.py | 3 ++- .../test_speculative_sampling.py | 1 - .../speculative_sequence/speculative_sequence.py | 12 ++++++------ .../target_hidden_projection.py | 2 +- 8 files changed, 17 insertions(+), 23 deletions(-) diff --git a/tests/turbomind/attention/test_verification_attention.py b/tests/turbomind/attention/test_verification_attention.py index def78ed965..8eda17a58a 100644 --- a/tests/turbomind/attention/test_verification_attention.py +++ b/tests/turbomind/attention/test_verification_attention.py @@ -3,7 +3,7 @@ import pytest import torch -import _turbomind as _tm +from lmdeploy.turbomind import _tm from .verification_attention import run_verification_attention @@ -84,13 +84,6 @@ def _reference(case, prefix_k, prefix_v, packed_qkv, q_bias, for request, history in enumerate(case['histories']): prefix_end = prefix_begin + history query_end = query_begin + p - prefix_positions = _positions(request, - history, - mrope_mode=case['mrope_mode'], - position_ids=mrope_position_ids, - position_delta=mrope_position_delta, - mrope_length=mrope_length, - device=packed_qkv.device) all_positions = _positions(request, history + p, mrope_mode=case['mrope_mode'], diff --git a/tests/turbomind/attention/verification_attention.py b/tests/turbomind/attention/verification_attention.py index 93db6fa3d4..525123a3e6 100644 --- a/tests/turbomind/attention/verification_attention.py +++ b/tests/turbomind/attention/verification_attention.py @@ -1,6 +1,7 @@ -import _turbomind as _tm import torch +from lmdeploy.turbomind import _tm + def run_verification_attention( *, diff --git a/tests/turbomind/draft_carry/draft_carry.py b/tests/turbomind/draft_carry/draft_carry.py index b4fe00a337..1cffd81528 100644 --- a/tests/turbomind/draft_carry/draft_carry.py +++ b/tests/turbomind/draft_carry/draft_carry.py @@ -7,7 +7,7 @@ def _load_native_bridge(): try: - import _turbomind as tm + from lmdeploy.turbomind import _turbomind as tm except ImportError: return None return tm if hasattr(tm, _NATIVE_SYMBOL) else None diff --git a/tests/turbomind/draft_carry/test_draft_carry.py b/tests/turbomind/draft_carry/test_draft_carry.py index 2deee02527..01bc127b8b 100644 --- a/tests/turbomind/draft_carry/test_draft_carry.py +++ b/tests/turbomind/draft_carry/test_draft_carry.py @@ -5,15 +5,15 @@ import pytest import torch +from .draft_carry import ( + is_available, + select_draft_carry, +) from .reference import ( OwnedTokenRows, compute_token_ownership, select_draft_carry_reference, ) -from .draft_carry import ( - is_available, - select_draft_carry, -) pytestmark = pytest.mark.skipif( not torch.cuda.is_available(), diff --git a/tests/turbomind/speculative_sampling/speculative_sampling.py b/tests/turbomind/speculative_sampling/speculative_sampling.py index 78e9a5a7e5..222da2efe8 100644 --- a/tests/turbomind/speculative_sampling/speculative_sampling.py +++ b/tests/turbomind/speculative_sampling/speculative_sampling.py @@ -1,6 +1,7 @@ -import _turbomind as _tm import torch +from lmdeploy.turbomind import _tm + STATE_CAPACITY_PER_LOGICAL_STATE = 4096 diff --git a/tests/turbomind/speculative_sampling/test_speculative_sampling.py b/tests/turbomind/speculative_sampling/test_speculative_sampling.py index 26cac21da9..97aec2d410 100644 --- a/tests/turbomind/speculative_sampling/test_speculative_sampling.py +++ b/tests/turbomind/speculative_sampling/test_speculative_sampling.py @@ -3,7 +3,6 @@ import pytest import torch -from .reference import greedy_reference, recovery_token_reference from .speculative_sampling import ( allocate_random_states, append_one_token_and_advance_sequence, diff --git a/tests/turbomind/speculative_sequence/speculative_sequence.py b/tests/turbomind/speculative_sequence/speculative_sequence.py index c1f9da48b5..504cb51e34 100644 --- a/tests/turbomind/speculative_sequence/speculative_sequence.py +++ b/tests/turbomind/speculative_sequence/speculative_sequence.py @@ -14,7 +14,7 @@ def initialize_target_verification( accepted_draft_count: torch.Tensor | None, request_to_generation_row_offsets: torch.Tensor, ) -> None: - import _turbomind + from lmdeploy.turbomind import _turbomind stream = torch.cuda.current_stream(entry_sequence_length.device) @@ -39,7 +39,7 @@ def build_draft_extension_key_offsets( accept_len: torch.Tensor, extension_index: int, ) -> None: - import _turbomind + from lmdeploy.turbomind import _turbomind stream = torch.cuda.current_stream(q_offsets.device) @@ -60,7 +60,7 @@ def stop_criteria( sequence_length_limit: torch.Tensor, finished: torch.Tensor, ) -> None: - import _turbomind + from lmdeploy.turbomind import _turbomind stream = torch.cuda.current_stream(token_ids_ptrs.device) @@ -82,7 +82,7 @@ def speculative_stop_criteria( sequence_length_limit: torch.Tensor, finished: torch.Tensor, ) -> None: - import _turbomind + from lmdeploy.turbomind import _turbomind stream = torch.cuda.current_stream(token_ids_ptrs.device) @@ -109,7 +109,7 @@ def build_draft_refresh_inputs( limit_to_accept_len: torch.Tensor, finished: torch.Tensor, ) -> None: - import _turbomind + from lmdeploy.turbomind import _turbomind stream = torch.cuda.current_stream(refresh_q_offsets.device) @@ -139,7 +139,7 @@ def draft_argmax_and_store_token( proposal_index: int, vocab_size: int, ) -> None: - import _turbomind + from lmdeploy.turbomind import _turbomind stream = torch.cuda.current_stream(logits.device) diff --git a/tests/turbomind/target_hidden_projection/target_hidden_projection.py b/tests/turbomind/target_hidden_projection/target_hidden_projection.py index e073eb3be8..2f19edda4f 100644 --- a/tests/turbomind/target_hidden_projection/target_hidden_projection.py +++ b/tests/turbomind/target_hidden_projection/target_hidden_projection.py @@ -7,7 +7,7 @@ def _load_native_bridge(): try: - import _turbomind as tm + from lmdeploy.turbomind import _turbomind as tm except ImportError: return None return tm if hasattr(tm, _NATIVE_SYMBOL) else None