From e99b5e100ef4a077d006120155c681f861c9c51a Mon Sep 17 00:00:00 2001 From: illu Date: Fri, 10 Jul 2026 09:58:17 +0000 Subject: [PATCH 01/20] Add full FT development snapshot --- kt-kernel/ext_bindings.cpp | 27 ++++- kt-kernel/operators/amx/sft_moe.hpp | 98 +++++++++++++++- kt-kernel/operators/common.hpp | 9 ++ kt-kernel/operators/moe-sft-tp.hpp | 29 ++++- kt-kernel/python/experts.py | 6 + kt-kernel/python/sft/__init__.py | 13 ++- kt-kernel/python/sft/amx.py | 166 +++++++++++++++++++++++----- kt-kernel/python/sft/arch.py | 14 ++- kt-kernel/python/sft/autograd.py | 67 ++++++++--- kt-kernel/python/sft/base.py | 115 ++++++++++++++----- kt-kernel/python/sft/config.py | 8 ++ kt-kernel/python/sft/layer.py | 126 +++++++++++++++++---- kt-kernel/python/sft/lora.py | 160 +++++++++++++++++++++++---- kt-kernel/python/sft/weights.py | 38 +++++-- kt-kernel/python/sft/wrapper.py | 88 ++++++++++++--- 15 files changed, 809 insertions(+), 155 deletions(-) diff --git a/kt-kernel/ext_bindings.cpp b/kt-kernel/ext_bindings.cpp index 87074e417..95557973e 100644 --- a/kt-kernel/ext_bindings.cpp +++ b/kt-kernel/ext_bindings.cpp @@ -325,22 +325,29 @@ class MOESFTBindings { intptr_t grad_down_lora_a; intptr_t grad_down_lora_b; intptr_t grad_weights; + intptr_t grad_gate_proj; + intptr_t grad_up_proj; + intptr_t grad_down_proj; }; static void inner(void* args) { Args* args_ = (Args*)args; args_->cpuinfer->enqueue(&TP_MOE_SFT::backward_binding, args_->moe, args_->grad_output, args_->grad_input, args_->grad_gate_lora_a, args_->grad_gate_lora_b, args_->grad_up_lora_a, args_->grad_up_lora_b, args_->grad_down_lora_a, args_->grad_down_lora_b, - args_->grad_weights); + args_->grad_weights, args_->grad_gate_proj, args_->grad_up_proj, + args_->grad_down_proj); } static std::pair cpuinfer_interface(std::shared_ptr> moe, intptr_t grad_output, intptr_t grad_input, intptr_t grad_gate_lora_a, intptr_t grad_gate_lora_b, intptr_t grad_up_lora_a, intptr_t grad_up_lora_b, intptr_t grad_down_lora_a, - intptr_t grad_down_lora_b, intptr_t grad_weights) { + intptr_t grad_down_lora_b, intptr_t grad_weights, + intptr_t grad_gate_proj, intptr_t grad_up_proj, + intptr_t grad_down_proj) { Args* args = new Args{nullptr, moe.get(), grad_output, grad_input, grad_gate_lora_a, grad_gate_lora_b, grad_up_lora_a, grad_up_lora_b, - grad_down_lora_a, grad_down_lora_b, grad_weights}; + grad_down_lora_a, grad_down_lora_b, grad_weights, + grad_gate_proj, grad_up_proj, grad_down_proj}; return std::make_pair((intptr_t)&inner, (intptr_t)args); } }; @@ -403,7 +410,13 @@ void bind_moe_sft_module(py::module_& moe_module, const char* name) { self.prepare_and_save_bwd((void*)gate, (void*)up, (void*)down, path); }) .def("submit_backward_repack", &MoeClass::submit_backward_repack) - .def("wait_backward_repack", &MoeClass::wait_backward_repack); + .def("wait_backward_repack", &MoeClass::wait_backward_repack) + // Update base weight BF16 pointers for reload_base_weights (full mode training) + // After calling this, call load_weights_task() to re-quantize BF16->AMX + .def("set_base_weight_pointers", + [](MoeClass& self, intptr_t gate, intptr_t up, intptr_t down) { + self.set_base_weight_pointers((void*)gate, (void*)up, (void*)down); + }); } #endif // defined(__x86_64__) && defined(USE_AMX_AVX_KERNEL) @@ -779,7 +792,11 @@ PYBIND11_MODULE(kt_kernel_ext, m) { .DEF_PTR_PROPERTY(MOESFTConfig, up_lora_a) .DEF_PTR_PROPERTY(MOESFTConfig, up_lora_b) .DEF_PTR_PROPERTY(MOESFTConfig, down_lora_a) - .DEF_PTR_PROPERTY(MOESFTConfig, down_lora_b); + .DEF_PTR_PROPERTY(MOESFTConfig, down_lora_b) + .def_readwrite("full_weight_grad", &MOESFTConfig::full_weight_grad) + .DEF_PTR_PROPERTY(MOESFTConfig, grad_gate_proj) + .DEF_PTR_PROPERTY(MOESFTConfig, grad_up_proj) + .DEF_PTR_PROPERTY(MOESFTConfig, grad_down_proj); py::class_>(moe_module, "MoE_Interface"); diff --git a/kt-kernel/operators/amx/sft_moe.hpp b/kt-kernel/operators/amx/sft_moe.hpp index 357ad9384..b516bcf75 100644 --- a/kt-kernel/operators/amx/sft_moe.hpp +++ b/kt-kernel/operators/amx/sft_moe.hpp @@ -1394,7 +1394,8 @@ class AMX_SFT_MOE_TP : public BaseMOE { void backward(const void* grad_output, void* grad_input, void* grad_gate_lora_a, void* grad_gate_lora_b, void* grad_up_lora_a, void* grad_up_lora_b, void* grad_down_lora_a, void* grad_down_lora_b, void* grad_weights, int full_intermediate_size = 0, float* fp32_grad_down_lora_b = nullptr, - float* fp32_grad_gate_lora_a = nullptr, float* fp32_grad_up_lora_a = nullptr) { + float* fp32_grad_gate_lora_a = nullptr, float* fp32_grad_up_lora_a = nullptr, + void* grad_gate_proj = nullptr, void* grad_up_proj = nullptr, void* grad_down_proj = nullptr) { // If full_intermediate_size not provided, use local (non-TP mode) if (full_intermediate_size == 0) full_intermediate_size = config_.intermediate_size; SFT_POOL_LOG("bwd_enter", config_.layer_idx, tp_part_idx, 0, cache_stack_top_, forward_pool_bytes_, @@ -1878,7 +1879,14 @@ class AMX_SFT_MOE_TP : public BaseMOE { print_grad_stats_fp32("grad_weights", (const float*)grad_weights, qlen * k); } - // ★ Cache pool is NOT freed here — kept for reuse across steps. + // ===================================================================== + // Step 5: Base weight gradient accumulation (full weight grad mode) + // ===================================================================== + if (sft_config_.full_weight_grad && grad_gate_proj && grad_up_proj && grad_down_proj) { + backward_base_weight_grad(cache, grad_output, grad_gate_proj, grad_up_proj, grad_down_proj); + } + + // \u2605 Cache pool is NOT freed here \u2014 kept for reuse across steps. // alloc_or_resize_cache_pool() is grow-only, so same-seqlen steps // reuse the existing allocation without malloc/free overhead. // Previously: free_seqlen_buffers() was called here, costing ~3.6ms per TP. @@ -1887,6 +1895,92 @@ class AMX_SFT_MOE_TP : public BaseMOE { cache.valid = false; } + /** + * @brief Compute base weight gradients via outer-product accumulation. + * + * For each activated expert e, computes: + * grad_gate_proj[e] = grad_gate_out[e]^T @ input[e] -> [I, H] + * grad_up_proj[e] = grad_up_out[e]^T @ input[e] -> [I, H] + * grad_down_proj[e] = grad_output[e]^T @ intermediate[e] -> [H, I] + * + * Uses FP32 accumulator for precision, writes BF16 output. + */ + void backward_base_weight_grad(const ForwardCache& cache, const void* grad_output, void* grad_gate_proj, + void* grad_up_proj, void* grad_down_proj) { + const int H = config_.hidden_size; + const int I = config_.intermediate_size; + const int E = config_.expert_num; + int activated_expert = cache.activated_expert_cache; + + auto* ggp = static_cast(grad_gate_proj); // [E, I, H] + auto* gup_ptr = static_cast(grad_up_proj); // [E, I, H] + auto* gdp = static_cast(grad_down_proj); // [E, H, I] + auto* grad_out_bf16 = static_cast(grad_output); + + for (int task_id = 0; task_id < activated_expert; task_id++) { + int expert_idx = cache.m_expert_id_map_cache[task_id]; + int m = cache.m_local_num_cache[expert_idx]; + if (m == 0) continue; + + int pos_start = 0; + for (int prev_id = 0; prev_id < task_id; prev_id++) { + pos_start += cache.m_local_num_cache[cache.m_expert_id_map_cache[prev_id]]; + } + + const auto& local_pos = cache.m_local_pos_cache[expert_idx]; + + // Allocate FP32 accumulators from forward pool (safe during backward) + float* acc_gate = static_cast(forward_pool_); // [I, H] + float* acc_up = acc_gate + (size_t)I * H; // [I, H] + float* acc_down = acc_up + (size_t)I * H; // [H, I] + + std::memset(acc_gate, 0, (size_t)I * H * sizeof(float)); + std::memset(acc_up, 0, (size_t)I * H * sizeof(float)); + std::memset(acc_down, 0, (size_t)H * I * sizeof(float)); + + for (int t = 0; t < m; t++) { + int tok_pos = local_pos[t]; + const ggml_bf16_t* input_row = cache.input_cache + (size_t)tok_pos * H; + const ggml_bf16_t* gate_grad_row = grad_gate_output_ + (size_t)(pos_start + t) * I; + const ggml_bf16_t* up_grad_row = grad_up_output_ + (size_t)(pos_start + t) * I; + const ggml_bf16_t* inter_row = cache.intermediate_cache + (size_t)(pos_start + t) * I; + const ggml_bf16_t* grad_out_row = grad_out_bf16 + (size_t)tok_pos * H; + + // gate_proj grad: [I, H] += grad_gate_out[t]^T @ input[t] + for (int i = 0; i < I; i++) { + float gg = GGML_BF16_TO_FP32(gate_grad_row[i]); + float gu = GGML_BF16_TO_FP32(up_grad_row[i]); + for (int h = 0; h < H; h++) { + float inp = GGML_BF16_TO_FP32(input_row[h]); + acc_gate[i * H + h] += gg * inp; + acc_up[i * H + h] += gu * inp; + } + } + + // down_proj grad: [H, I] += grad_output[t]^T @ intermediate[t] + for (int h = 0; h < H; h++) { + float go = GGML_BF16_TO_FP32(grad_out_row[h]); + for (int i = 0; i < I; i++) { + acc_down[h * I + i] += go * GGML_BF16_TO_FP32(inter_row[i]); + } + } + } + + // Convert FP32 accumulators to BF16 and store + for (int i = 0; i < I; i++) { + for (int h = 0; h < H; h++) { + ggp[(size_t)expert_idx * I * H + (size_t)i * H + h] = GGML_FP32_TO_BF16(acc_gate[i * H + h]); + gup_ptr[(size_t)expert_idx * I * H + (size_t)i * H + h] = GGML_FP32_TO_BF16(acc_up[i * H + h]); + } + } + for (int h = 0; h < H; h++) { + for (int i = 0; i < I; i++) { + gdp[(size_t)expert_idx * H * I + (size_t)h * I + i] = GGML_FP32_TO_BF16(acc_down[h * I + i]); + } + } + } + } + /** * @brief Get qlen from the top of the forward cache stack. * diff --git a/kt-kernel/operators/common.hpp b/kt-kernel/operators/common.hpp index 86f3464c6..5c39c00d8 100644 --- a/kt-kernel/operators/common.hpp +++ b/kt-kernel/operators/common.hpp @@ -351,6 +351,15 @@ struct MOESFTConfig : public GeneralMOEConfig { void* down_lora_a = nullptr; // [expert_num, lora_rank, intermediate_size] void* down_lora_b = nullptr; // [expert_num, hidden_size, lora_rank] + // Full weight gradient configuration + bool full_weight_grad = false; + + // Base weight gradient buffer pointers (directly pointing to Python tensor memory, zero-copy) + // Only used when full_weight_grad == true + void* grad_gate_proj = nullptr; // [expert_num, intermediate_size, hidden_size] + void* grad_up_proj = nullptr; // [expert_num, intermediate_size, hidden_size] + void* grad_down_proj = nullptr; // [expert_num, hidden_size, intermediate_size] + MOESFTConfig() : GeneralMOEConfig() {} MOESFTConfig(int expert_num, int routed_expert_num, int hidden_size, int intermediate_size) diff --git a/kt-kernel/operators/moe-sft-tp.hpp b/kt-kernel/operators/moe-sft-tp.hpp index 4e8da6ad0..c8a0e0986 100644 --- a/kt-kernel/operators/moe-sft-tp.hpp +++ b/kt-kernel/operators/moe-sft-tp.hpp @@ -345,7 +345,7 @@ class TP_MOE_SFT : public TP_MOE { throw std::runtime_error("K2 pre-quantized mode does not support TP > 1 yet"); } } else if (config.gate_proj != nullptr) { - printf("TP_MOE_SFT: From BF16 with partitioning\n"); + // printf("TP_MOE_SFT: From BF16 with partitioning\n"); // Temporary storage for partitioned weights std::vector temp_gate(tp_count); @@ -549,7 +549,8 @@ class TP_MOE_SFT : public TP_MOE { */ void backward(const void* grad_output, void* grad_input, void* grad_gate_lora_a, void* grad_gate_lora_b, void* grad_up_lora_a, void* grad_up_lora_b, void* grad_down_lora_a, void* grad_down_lora_b, - void* grad_weights) { + void* grad_weights, void* grad_gate_proj = nullptr, void* grad_up_proj = nullptr, + void* grad_down_proj = nullptr) { auto pool = config.pool; // Get full intermediate_size (before TP partitioning) @@ -693,7 +694,8 @@ class TP_MOE_SFT : public TP_MOE { tp_down_a_ptr[numa_id], /* copy-type: direct write */ nullptr, /* grad_down_lora_b — unused, FP32 path below */ part_grad_weights_[numa_id], full_intermediate_size, tp_fp32_down_b[numa_id], - tp_fp32_gate_a[numa_id], tp_fp32_up_a[numa_id]); + tp_fp32_gate_a[numa_id], tp_fp32_up_a[numa_id], grad_gate_proj, grad_up_proj, + grad_down_proj); }); // // Collect per-thread timing from all NUMA subpools @@ -889,10 +891,11 @@ class TP_MOE_SFT : public TP_MOE { */ void backward_binding(intptr_t grad_output, intptr_t grad_input, intptr_t grad_gate_lora_a, intptr_t grad_gate_lora_b, intptr_t grad_up_lora_a, intptr_t grad_up_lora_b, intptr_t grad_down_lora_a, - intptr_t grad_down_lora_b, intptr_t grad_weights) { + intptr_t grad_down_lora_b, intptr_t grad_weights, intptr_t grad_gate_proj, + intptr_t grad_up_proj, intptr_t grad_down_proj) { backward((const void*)grad_output, (void*)grad_input, (void*)grad_gate_lora_a, (void*)grad_gate_lora_b, (void*)grad_up_lora_a, (void*)grad_up_lora_b, (void*)grad_down_lora_a, (void*)grad_down_lora_b, - (void*)grad_weights); + (void*)grad_weights, (void*)grad_gate_proj, (void*)grad_up_proj, (void*)grad_down_proj); } /** @@ -1099,6 +1102,22 @@ class TP_MOE_SFT : public TP_MOE { update_lora_weights((void*)gate_lora_a, (void*)gate_lora_b, (void*)up_lora_a, (void*)up_lora_b, (void*)down_lora_a, (void*)down_lora_b); } + + /** + * @brief Update base weight BF16 pointers for reload_base_weights (full mode training). + * + * After calling this, call load_weights_task() to re-quantize BF16->AMX + * and update the C++ kernel's internal quantized buffers. + * This avoids creating a new C++ MOE object (~0.6s/layer for quantization + * vs ~1.9s/layer for full object recreation). + */ + void set_base_weight_pointers(void* gate, void* up, void* down) { + config.gate_proj = gate; + config.up_proj = up; + config.down_proj = down; + // Mark that weights need re-loading (partitioning + quantization) + weights_loaded = false; + } }; #endif // CPUINFER_OPERATOR_MOE_SFT_TP_HPP diff --git a/kt-kernel/python/experts.py b/kt-kernel/python/experts.py index 318fe3170..d31e42801 100644 --- a/kt-kernel/python/experts.py +++ b/kt-kernel/python/experts.py @@ -144,6 +144,8 @@ def __new__( # Quantization config (for K-Group SFT methods) group_size: int = 128, zero_point: bool = True, + # Full weight gradient mode (for full fine-tuning without LoRA) + full_weight_grad: bool = False, # V4-Flash 2604B SwiGLU clamp limit. 0.0 = disabled (default for # every dtype except DSV4-2604B routed experts, which set this to # 10.0 to match trtllm gemm1_clamp_limit / deep_gemm @@ -252,6 +254,7 @@ def __new__( max_cache_depth=max_cache_depth, group_size=group_size, zero_point=zero_point, + full_weight_grad=full_weight_grad, ) # Forward static methods to the base class @@ -297,6 +300,7 @@ def clear_sft_buffer_cache(): to reset the buffer state or free memory during SFT. """ from .sft.base import KExpertsSFTBuffer + KExpertsSFTBuffer.clear_cache() @@ -400,6 +404,7 @@ def _create_sft_wrapper( max_cache_depth: int, group_size: int, zero_point: bool, + full_weight_grad: bool = False, ): """ Create an SFT wrapper based on the method. @@ -430,4 +435,5 @@ def _create_sft_wrapper( method=method, group_size=group_size, zero_point=zero_point, + full_weight_grad=full_weight_grad, ) diff --git a/kt-kernel/python/sft/__init__.py b/kt-kernel/python/sft/__init__.py index 7cab43bd2..88b266e08 100644 --- a/kt-kernel/python/sft/__init__.py +++ b/kt-kernel/python/sft/__init__.py @@ -14,8 +14,15 @@ from .base import BaseSFTMoEWrapper, KExpertsSFTBuffer from .amx import AMXSFTMoEWrapper from .arch import ( - MOEArchConfig, get_moe_arch_config, get_moe_module, move_non_experts_to_gpu, get_expert_device, - KTAMXError, KTAMXNotAvailableError, KTAMXModelNotSupportedError, KTAMXConfigError, + MOEArchConfig, + get_moe_arch_config, + get_moe_module, + move_non_experts_to_gpu, + get_expert_device, + KTAMXError, + KTAMXNotAvailableError, + KTAMXModelNotSupportedError, + KTAMXConfigError, ) from .autograd import KTMoEFunction from .layer import KTMoELayerWrapper @@ -28,6 +35,7 @@ from .lora import ( kt_adapt_peft_lora, get_kt_lora_params, + get_kt_trainable_params, update_kt_lora_pointers, sync_kt_lora_gradients, save_lora_experts_to_adapter, @@ -67,6 +75,7 @@ "INT8ExpertWeights", "kt_adapt_peft_lora", "get_kt_lora_params", + "get_kt_trainable_params", "update_kt_lora_pointers", "sync_kt_lora_gradients", "save_lora_experts_to_adapter", diff --git a/kt-kernel/python/sft/amx.py b/kt-kernel/python/sft/amx.py index effa7dfd2..c59a8b455 100644 --- a/kt-kernel/python/sft/amx.py +++ b/kt-kernel/python/sft/amx.py @@ -9,6 +9,7 @@ from __future__ import annotations import ctypes +import logging import os import glob as _glob import torch @@ -16,6 +17,8 @@ from kt_kernel_ext.moe import MOESFTConfig +logger = logging.getLogger(__name__) + from ..utils.loader import BF16SafeTensorLoader, SafeTensorLoader try: @@ -81,6 +84,7 @@ def __init__( method: str = "AMXBF16_SFT", group_size: int = 128, zero_point: bool = True, + full_weight_grad: bool = False, ): if not _HAS_AMX_SFT_SUPPORT: raise RuntimeError( @@ -102,6 +106,7 @@ def __init__( lora_rank=lora_rank, lora_alpha=lora_alpha, max_cache_depth=max_cache_depth, + full_weight_grad=full_weight_grad, ) self.method = method @@ -140,19 +145,42 @@ def _make_backward_task(self, buffer: KExpertsSFTBuffer): return self.moe.backward_task( buffer.grad_output_cpu.data_ptr(), buffer.grad_input_cpu.data_ptr(), - 0, 0, 0, 0, 0, 0, + 0, + 0, + 0, + 0, + 0, + 0, buffer.grad_weights.data_ptr(), + 0, + 0, + 0, # grad_gate_proj, grad_up_proj, grad_down_proj ) + + # Base weight grad pointers (nullptr if not in full mode) + grad_gate_proj_ptr = ( + self.grad_gate_proj_buf.data_ptr() if self._full_weight_grad and self.grad_gate_proj_buf is not None else 0 + ) + grad_up_proj_ptr = ( + self.grad_up_proj_buf.data_ptr() if self._full_weight_grad and self.grad_up_proj_buf is not None else 0 + ) + grad_down_proj_ptr = ( + self.grad_down_proj_buf.data_ptr() if self._full_weight_grad and self.grad_down_proj_buf is not None else 0 + ) + return self.moe.backward_task( buffer.grad_output_cpu.data_ptr(), buffer.grad_input_cpu.data_ptr(), - self.grad_gate_lora_a.data_ptr(), - self.grad_gate_lora_b.data_ptr(), - self.grad_up_lora_a.data_ptr(), - self.grad_up_lora_b.data_ptr(), - self.grad_down_lora_a.data_ptr(), - self.grad_down_lora_b.data_ptr(), + self.grad_gate_lora_a.data_ptr() if self.lora_rank > 0 else 0, + self.grad_gate_lora_b.data_ptr() if self.lora_rank > 0 else 0, + self.grad_up_lora_a.data_ptr() if self.lora_rank > 0 else 0, + self.grad_up_lora_b.data_ptr() if self.lora_rank > 0 else 0, + self.grad_down_lora_a.data_ptr() if self.lora_rank > 0 else 0, + self.grad_down_lora_b.data_ptr() if self.lora_rank > 0 else 0, buffer.grad_weights.data_ptr(), + grad_gate_proj_ptr, + grad_up_proj_ptr, + grad_down_proj_ptr, ) # ========== Weight loading ========== @@ -187,6 +215,7 @@ def load_weights(self, physical_to_logical_map_cpu: torch.Tensor) -> None: config.layer_idx = self.layer_idx config.share_backward_bb = getattr(self, "share_backward_bb", False) config.share_cache_pool = getattr(self, "share_cache_pool", False) + config.full_weight_grad = self._full_weight_grad config.physical_to_logical_map = self._physical_to_logical_map_cpu.data_ptr() if getattr(self, "_use_kt_direct_load", False): @@ -229,6 +258,13 @@ def load_weights(self, physical_to_logical_map_cpu: torch.Tensor) -> None: config.quant_config.group_size = self.group_size config.quant_config.zero_point = self.zero_point + # Release old C++ MOE object before creating a new one to avoid memory leak + old_moe = getattr(self, "moe", None) + if old_moe is not None: + del old_moe + import gc + gc.collect() + self.moe = self._moe_class(config) self.cpu_infer.submit(self.moe.load_weights_task()) @@ -239,9 +275,11 @@ def load_weights(self, physical_to_logical_map_cpu: torch.Tensor) -> None: self.cpu_infer.sync() # Release Python-side weight tensors (C++ copied them) - self.gate_proj = None - self.up_proj = None - self.down_proj = None + # In full_weight_grad mode, keep them for nn.Parameter initialization + if not self._full_weight_grad: + self.gate_proj = None + self.up_proj = None + self.down_proj = None if getattr(self, "_bf16_gate_proj", None) is not None: self._bf16_gate_proj = None @@ -250,18 +288,34 @@ def load_weights(self, physical_to_logical_map_cpu: torch.Tensor) -> None: if getattr(self, "_use_projs_path", False): for attr in [ - "_gate_weights_per_numa", "_up_weights_per_numa", "_down_weights_per_numa", - "_gate_scales_per_numa", "_up_scales_per_numa", "_down_scales_per_numa", - "_gate_projs_ptrs", "_up_projs_ptrs", "_down_projs_ptrs", - "_gate_scale_ptrs", "_up_scale_ptrs", "_down_scale_ptrs", + "_gate_weights_per_numa", + "_up_weights_per_numa", + "_down_weights_per_numa", + "_gate_scales_per_numa", + "_up_scales_per_numa", + "_down_scales_per_numa", + "_gate_projs_ptrs", + "_up_projs_ptrs", + "_down_projs_ptrs", + "_gate_scale_ptrs", + "_up_scale_ptrs", + "_down_scale_ptrs", ]: setattr(self, attr, None) if getattr(self, "_has_bwd_projs", False): for attr in [ - "_gate_bwd_weights_per_numa", "_up_bwd_weights_per_numa", "_down_bwd_weights_per_numa", - "_gate_bwd_scales_per_numa", "_up_bwd_scales_per_numa", "_down_bwd_scales_per_numa", - "_gate_bwd_projs_ptrs", "_up_bwd_projs_ptrs", "_down_bwd_projs_ptrs", - "_gate_bwd_scale_ptrs", "_up_bwd_scale_ptrs", "_down_bwd_scale_ptrs", + "_gate_bwd_weights_per_numa", + "_up_bwd_weights_per_numa", + "_down_bwd_weights_per_numa", + "_gate_bwd_scales_per_numa", + "_up_bwd_scales_per_numa", + "_down_bwd_scales_per_numa", + "_gate_bwd_projs_ptrs", + "_up_bwd_projs_ptrs", + "_down_bwd_projs_ptrs", + "_gate_bwd_scale_ptrs", + "_up_bwd_scale_ptrs", + "_down_bwd_scale_ptrs", ]: setattr(self, attr, None) @@ -312,6 +366,7 @@ def _load_base_weights_from_file(self) -> None: self.up_proj = torch.stack(up_weights, dim=0).contiguous() self.down_proj = torch.stack(down_weights, dim=0).contiguous() else: + def _make_ptrs(arrays_per_numa): return [ [ @@ -416,12 +471,18 @@ def _validate_prepartitioned_weights(self) -> None: def init_lora_weights( self, - gate_lora_a: torch.Tensor, gate_lora_b: torch.Tensor, - up_lora_a: torch.Tensor, up_lora_b: torch.Tensor, - down_lora_a: torch.Tensor, down_lora_b: torch.Tensor, - grad_gate_lora_a: torch.Tensor, grad_gate_lora_b: torch.Tensor, - grad_up_lora_a: torch.Tensor, grad_up_lora_b: torch.Tensor, - grad_down_lora_a: torch.Tensor, grad_down_lora_b: torch.Tensor, + gate_lora_a: torch.Tensor, + gate_lora_b: torch.Tensor, + up_lora_a: torch.Tensor, + up_lora_b: torch.Tensor, + down_lora_a: torch.Tensor, + down_lora_b: torch.Tensor, + grad_gate_lora_a: torch.Tensor, + grad_gate_lora_b: torch.Tensor, + grad_up_lora_a: torch.Tensor, + grad_up_lora_b: torch.Tensor, + grad_down_lora_a: torch.Tensor, + grad_down_lora_b: torch.Tensor, ) -> None: expected_shapes = { "gate_lora_a": (self.num_experts, self.lora_rank, self.hidden_size), @@ -432,9 +493,12 @@ def init_lora_weights( "down_lora_b": (self.num_experts, self.hidden_size, self.lora_rank), } provided = { - "gate_lora_a": gate_lora_a, "gate_lora_b": gate_lora_b, - "up_lora_a": up_lora_a, "up_lora_b": up_lora_b, - "down_lora_a": down_lora_a, "down_lora_b": down_lora_b, + "gate_lora_a": gate_lora_a, + "gate_lora_b": gate_lora_b, + "up_lora_a": up_lora_a, + "up_lora_b": up_lora_b, + "down_lora_a": down_lora_a, + "down_lora_b": down_lora_b, } for name, tensor in provided.items(): expected = expected_shapes[name] @@ -470,6 +534,8 @@ def update_lora_weights(self) -> None: if self._is_skip_lora: return if not self._lora_initialized: + if self.lora_rank <= 0: + return # Full mode without LoRA — no LoRA weights to update raise RuntimeError("LoRA weights not initialized. Call init_lora_weights() first.") # Weight pointer updates are load-time synchronous work. Calling the @@ -484,6 +550,52 @@ def update_lora_weights(self) -> None: self.down_lora_b.data_ptr(), ) + def update_base_weights(self) -> None: + """Sync updated base weight parameters back to C++ kernel after optimizer step.""" + if not self._weights_loaded: + raise RuntimeError("Weights not loaded. Call load_weights() first.") + if not self._full_weight_grad: + return # No base weights to update in LoRA mode + if self.gate_proj_buf is None: + raise RuntimeError("Base weight buffers not initialized. Call init_full_weight_grad_buffers() first.") + + logger.info(f"Layer {self.layer_idx}: update_base_weights() - syncing updated weights to C++ kernel") + + # Preferred path: update config pointers on existing C++ object and re-quantize. + # This avoids full C++ MOE object recreation (~0.6s/layer vs ~1.9s/layer). + if hasattr(self.moe, "set_base_weight_pointers"): + self.moe.set_base_weight_pointers( + self.gate_proj_buf.data.data_ptr(), + self.up_proj_buf.data.data_ptr(), + self.down_proj_buf.data.data_ptr(), + ) + self.cpu_infer.submit(self.moe.load_weights_task()) + self.cpu_infer.sync() + logger.info(f"Layer {self.layer_idx}: update_base_weights() - re-quantized existing kernel") + return + + # Fallback: full reload path (creates new C++ MOE object) + # This is slower but works without C++ set_base_weight_pointers support. + logger.warning( + f"Layer {self.layer_idx}: set_base_weight_pointers not available, " + f"falling back to full C++ MOE object recreation" + ) + old_moe = getattr(self, "moe", None) + if old_moe is not None: + del old_moe + + self.gate_proj = self.gate_proj_buf.data + self.up_proj = self.up_proj_buf.data + self.down_proj = self.down_proj_buf.data + physical_to_logical_map = torch.arange(self.num_experts, dtype=torch.int64, device="cpu") + self._weights_loaded = False # Allow re-load + self.load_weights_from_tensors( + gate_proj=self.gate_proj, + up_proj=self.up_proj, + down_proj=self.down_proj, + physical_to_logical_map_cpu=physical_to_logical_map, + ) + def save_backward_weights_from_tensors( self, gate_proj: torch.Tensor, diff --git a/kt-kernel/python/sft/arch.py b/kt-kernel/python/sft/arch.py index 43b2e2cbc..234926d61 100644 --- a/kt-kernel/python/sft/arch.py +++ b/kt-kernel/python/sft/arch.py @@ -110,6 +110,18 @@ def get_moe_arch_config(config) -> MOEArchConfig: num_experts_per_tok=cfg.num_experts_per_tok, has_shared_experts=getattr(cfg, "shared_expert_intermediate_size", 0) > 0, ) + if "Glm4Moe" in arch: + return MOEArchConfig( + moe_layer_attr="mlp", + router_attr="gate", + experts_attr="experts", + weight_names=("gate_proj", "up_proj", "down_proj"), + expert_num=config.n_routed_experts, + intermediate_size=config.moe_intermediate_size, + num_experts_per_tok=config.num_experts_per_tok, + has_shared_experts=getattr(config, "n_shared_experts", 0) > 0, + router_type="glm4_moe_gate", + ) if "Mixtral" in arch: return MOEArchConfig( moe_layer_attr="block_sparse_moe", @@ -124,7 +136,7 @@ def get_moe_arch_config(config) -> MOEArchConfig: raise KTAMXModelNotSupportedError( f"Model architecture {arch} not supported for KT AMX. " - "Supported architectures: DeepseekV2, DeepseekV3, Qwen2Moe, Qwen3Moe, Qwen3_5Moe, Mixtral" + "Supported architectures: DeepseekV2, DeepseekV3, Qwen2Moe, Qwen3Moe, Qwen3_5Moe, Glm4Moe, Mixtral" ) diff --git a/kt-kernel/python/sft/autograd.py b/kt-kernel/python/sft/autograd.py index 0264e9de6..98c92ed97 100644 --- a/kt-kernel/python/sft/autograd.py +++ b/kt-kernel/python/sft/autograd.py @@ -38,12 +38,17 @@ def forward( training: bool, train_lora: bool, all_qlens: list[int] | tuple[int, ...] | None, + gate_proj_param: torch.Tensor | None = None, + up_proj_param: torch.Tensor | None = None, + down_proj_param: torch.Tensor | None = None, ) -> torch.Tensor: if _KT_SFT_DEBUG: logging.debug( "KTMoEFunction.forward: layer=%d training=%s train_lora=%s", - layer_idx, training, train_lora, + layer_idx, + training, + train_lora, ) original_device = hidden_states.device @@ -52,6 +57,7 @@ def forward( qlen = batch_size * seq_len import torch.distributed as dist + dist_on = dist.is_initialized() and dist.get_world_size() > 1 rank = dist.get_rank() if dist.is_initialized() else 0 world_size = dist.get_world_size() if dist_on else 1 @@ -65,13 +71,9 @@ def forward( else: all_qlens_list = [int(q) for q in all_qlens] if len(all_qlens_list) != world_size: - raise RuntimeError( - f"all_qlens length mismatch: got {len(all_qlens_list)}, expected {world_size}" - ) + raise RuntimeError(f"all_qlens length mismatch: got {len(all_qlens_list)}, expected {world_size}") if int(all_qlens_list[rank]) != qlen: - raise RuntimeError( - f"Rank {rank} qlen mismatch: local={qlen}, all_qlens[{rank}]={all_qlens_list[rank]}" - ) + raise RuntimeError(f"Rank {rank} qlen mismatch: local={qlen}, all_qlens[{rank}]={all_qlens_list[rank]}") total_qlen = sum(all_qlens_list) # Rank 0: sync CPU result and split by real lengths @@ -100,9 +102,7 @@ def forward( output = cpu_output.view(batch_size, seq_len, hidden_size).to(dtype=original_dtype) else: # Broadcast-only rank (no wrapper) - output = torch.empty( - batch_size, seq_len, hidden_size, device=original_device, dtype=original_dtype - ) + output = torch.empty(batch_size, seq_len, hidden_size, device=original_device, dtype=original_dtype) ctx.wrapper = wrapper ctx.hidden_size = hidden_size @@ -120,6 +120,11 @@ def forward( ctx.num_experts_per_tok = num_experts_per_tok ctx.layer_idx = layer_idx + # Store base weight param references for gradient flow in full mode + ctx.full_weight_grad = ( + wrapper is not None and getattr(wrapper, "_full_weight_grad", False) and gate_proj_param is not None + ) + # Save a sentinel tensor so non-reentrant checkpoint's saved_tensors # hooks can intercept it. When backward accesses ctx.saved_tensors, # the checkpoint unpack hook triggers a full recompute of the decoder @@ -135,7 +140,7 @@ def forward( @staticmethod def backward(ctx, grad_output: torch.Tensor): # Wait for any in-flight async repack before recompute forward uses the pool - if getattr(ctx.wrapper, 'share_backward_bb', False): + if getattr(ctx.wrapper, "share_backward_bb", False): ctx.wrapper.wait_backward_repack() # Access saved_tensors FIRST — under non-reentrant checkpoint this @@ -152,12 +157,15 @@ def backward(ctx, grad_output: torch.Tensor): num_experts_per_tok = ctx.num_experts_per_tok import torch.distributed as dist + rank = dist.get_rank() if dist.is_initialized() else 0 if _KT_SFT_DEBUG: logging.debug( "KTMoEFunction.backward: layer=%d dist_on=%s qlen=%d", - getattr(ctx, "layer_idx", -1), dist_on, qlen, + getattr(ctx, "layer_idx", -1), + dist_on, + qlen, ) if dist_on: @@ -243,12 +251,39 @@ def backward(ctx, grad_output: torch.Tensor): grad_weights = grad_weights.to(dtype=torch.bfloat16) else: # No wrapper, no dist — shouldn't happen in normal flow - grad_input = torch.zeros(batch_size, seq_len, hidden_size, device=ctx.original_device, dtype=ctx.original_dtype) + grad_input = torch.zeros( + batch_size, seq_len, hidden_size, device=ctx.original_device, dtype=ctx.original_dtype + ) grad_weights = torch.zeros(ctx.weights_shape, device=ctx.weights_device, dtype=ctx.weights_dtype) # Trigger async repack for next MoE layer in backward order - next_bwd = getattr(ctx.wrapper, '_next_backward_wrapper', None) - if next_bwd is not None and getattr(next_bwd, 'share_backward_bb', False): + next_bwd = getattr(ctx.wrapper, "_next_backward_wrapper", None) + if next_bwd is not None and getattr(next_bwd, "share_backward_bb", False): next_bwd.submit_backward_repack() - return grad_input, None, grad_weights, None, None, None, None, None, None, None, None + # Base weight gradients: return C++-written grad buffers in full mode, None otherwise + if ctx.full_weight_grad and ctx.wrapper is not None: + grad_gate_proj = ctx.wrapper.grad_gate_proj_buf + grad_up_proj = ctx.wrapper.grad_up_proj_buf + grad_down_proj = ctx.wrapper.grad_down_proj_buf + else: + grad_gate_proj = None + grad_up_proj = None + grad_down_proj = None + + return ( + grad_input, + None, + grad_weights, + None, + None, + None, + None, + None, + None, + None, + None, + grad_gate_proj, + grad_up_proj, + grad_down_proj, + ) diff --git a/kt-kernel/python/sft/base.py b/kt-kernel/python/sft/base.py index 57ba8e424..c04f4cc6f 100644 --- a/kt-kernel/python/sft/base.py +++ b/kt-kernel/python/sft/base.py @@ -145,6 +145,7 @@ def __init__( lora_rank: int = 16, lora_alpha: float = 32.0, max_cache_depth: int = 1, + full_weight_grad: bool = False, ): self.cpu_infer = self._get_cpu_infer(cpuinfer_threads, threadpool_count) @@ -154,7 +155,7 @@ def __init__( moe_intermediate_size=moe_intermediate_size, num_experts_per_tok=num_experts_per_tok, ) - self._validate_sft_config(lora_rank, lora_alpha, max_cache_depth) + self._validate_sft_config(lora_rank, lora_alpha, max_cache_depth, full_weight_grad=full_weight_grad) self.layer_idx = layer_idx self.num_experts = num_experts @@ -168,9 +169,11 @@ def __init__( self.lora_rank = lora_rank self.lora_alpha = lora_alpha - self.lora_scaling = lora_alpha / lora_rank + self.lora_scaling = lora_alpha / lora_rank if lora_rank > 0 else 0.0 self.max_cache_depth = max_cache_depth + self._full_weight_grad = full_weight_grad + self.gate_lora_a: Optional[torch.Tensor] = None self.gate_lora_b: Optional[torch.Tensor] = None self.up_lora_a: Optional[torch.Tensor] = None @@ -178,22 +181,75 @@ def __init__( self.down_lora_a: Optional[torch.Tensor] = None self.down_lora_b: Optional[torch.Tensor] = None + # Base weight parameters for full fine-tuning + self.gate_proj_buf: Optional[torch.Tensor] = None + self.up_proj_buf: Optional[torch.Tensor] = None + self.down_proj_buf: Optional[torch.Tensor] = None + self.grad_gate_proj_buf: Optional[torch.Tensor] = None + self.grad_up_proj_buf: Optional[torch.Tensor] = None + self.grad_down_proj_buf: Optional[torch.Tensor] = None + self._weights_loaded: bool = False self._lora_initialized: bool = False self._cache_depth: int = 0 self._is_skip_lora: bool = False + self._base_weights_dirty: bool = False self.moe = None @staticmethod - def _validate_sft_config(lora_rank: int, lora_alpha: float, max_cache_depth: int) -> None: - if lora_rank <= 0: - raise ValueError(f"lora_rank must be positive, got {lora_rank}") - if lora_alpha <= 0: + def _validate_sft_config( + lora_rank: int, lora_alpha: float, max_cache_depth: int, full_weight_grad: bool = False + ) -> None: + if not full_weight_grad and lora_rank <= 0: + raise ValueError( + f"lora_rank must be positive in LoRA mode, got {lora_rank}. " + "Set kt_train_mode='full' for full fine-tuning." + ) + if lora_rank > 0 and lora_alpha <= 0: raise ValueError(f"lora_alpha must be positive, got {lora_alpha}") if max_cache_depth <= 0: raise ValueError(f"max_cache_depth must be positive, got {max_cache_depth}") + # ========== Full weight grad methods ========== + + def init_full_weight_grad_buffers( + self, gate_proj: torch.Tensor, up_proj: torch.Tensor, down_proj: torch.Tensor + ) -> None: + """Initialize base weight nn.Parameter buffers and gradient buffers for full fine-tuning. + + Args: + gate_proj: [num_experts, intermediate_size, hidden_size] BF16 CPU tensor + up_proj: [num_experts, intermediate_size, hidden_size] BF16 CPU tensor + down_proj: [num_experts, hidden_size, intermediate_size] BF16 CPU tensor + """ + import torch.nn as nn + + dtype = torch.bfloat16 + E = self.num_experts + I = self.moe_intermediate_size + H = self.hidden_size + + # Create nn.Parameter buffers (optimizer-visible) + self.gate_proj_buf = nn.Parameter(gate_proj.to(dtype=dtype, device="cpu").contiguous(), requires_grad=True) + self.up_proj_buf = nn.Parameter(up_proj.to(dtype=dtype, device="cpu").contiguous(), requires_grad=True) + self.down_proj_buf = nn.Parameter(down_proj.to(dtype=dtype, device="cpu").contiguous(), requires_grad=True) + + # Create gradient buffers (C++ writes directly to these) + self.grad_gate_proj_buf = torch.zeros(E, I, H, dtype=dtype, device="cpu") + self.grad_up_proj_buf = torch.zeros(E, I, H, dtype=dtype, device="cpu") + self.grad_down_proj_buf = torch.zeros(E, H, I, dtype=dtype, device="cpu") + + # Note: .grad is NOT pre-assigned here. PyTorch autograd will set it + # when KTMoEFunction.backward() returns the gradient buffers. + # The C++ kernel writes directly to grad_gate_proj_buf etc., + # and backward returns them so PyTorch can propagate correctly. + + @abstractmethod + def update_base_weights(self) -> None: + """Sync updated base weight parameters back to C++ kernel after optimizer step.""" + ... + # ========== Abstract methods for subclasses ========== @abstractmethod @@ -207,24 +263,27 @@ def _make_backward_task(self, buffer: KExpertsSFTBuffer): ... @abstractmethod - def load_weights(self, physical_to_logical_map_cpu: torch.Tensor) -> None: - ... + def load_weights(self, physical_to_logical_map_cpu: torch.Tensor) -> None: ... @abstractmethod def init_lora_weights( self, - gate_lora_a: torch.Tensor, gate_lora_b: torch.Tensor, - up_lora_a: torch.Tensor, up_lora_b: torch.Tensor, - down_lora_a: torch.Tensor, down_lora_b: torch.Tensor, - grad_gate_lora_a: torch.Tensor, grad_gate_lora_b: torch.Tensor, - grad_up_lora_a: torch.Tensor, grad_up_lora_b: torch.Tensor, - grad_down_lora_a: torch.Tensor, grad_down_lora_b: torch.Tensor, - ) -> None: - ... + gate_lora_a: torch.Tensor, + gate_lora_b: torch.Tensor, + up_lora_a: torch.Tensor, + up_lora_b: torch.Tensor, + down_lora_a: torch.Tensor, + down_lora_b: torch.Tensor, + grad_gate_lora_a: torch.Tensor, + grad_gate_lora_b: torch.Tensor, + grad_up_lora_a: torch.Tensor, + grad_up_lora_b: torch.Tensor, + grad_down_lora_a: torch.Tensor, + grad_down_lora_b: torch.Tensor, + ) -> None: ... @abstractmethod - def update_lora_weights(self) -> None: - ... + def update_lora_weights(self) -> None: ... # ========== Buffer helpers ========== @@ -242,7 +301,7 @@ def _get_buffer(self, qlen: int) -> KExpertsSFTBuffer: def _validate_forward_inputs(self, hidden_states: torch.Tensor, expert_ids: torch.Tensor, weights: torch.Tensor): if not self._weights_loaded: raise RuntimeError("Weights not loaded. Call load_weights() or load_weights_from_tensors() first.") - if not self._lora_initialized and not self._is_skip_lora: + if not self._lora_initialized and not self._is_skip_lora and not self._full_weight_grad: raise RuntimeError("LoRA weights not initialized. Call init_lora_weights() first.") qlen = hidden_states.shape[0] if qlen > self.chunked_prefill_size: @@ -255,12 +314,16 @@ def _validate_forward_inputs(self, hidden_states: torch.Tensor, expert_ids: torc f"expert_ids shape {tuple(expert_ids.shape)} must be ({qlen}, {self.num_experts_per_tok})." ) if weights.shape[0] != qlen or weights.shape[1] != self.num_experts_per_tok: - raise ValueError( - f"weights shape {tuple(weights.shape)} must be ({qlen}, {self.num_experts_per_tok})." - ) + raise ValueError(f"weights shape {tuple(weights.shape)} must be ({qlen}, {self.num_experts_per_tok}).") - def _copy_inputs_to_buffer(self, buffer: KExpertsSFTBuffer, hidden_states: torch.Tensor, - expert_ids: torch.Tensor, weights: torch.Tensor, qlen: int) -> torch.device: + def _copy_inputs_to_buffer( + self, + buffer: KExpertsSFTBuffer, + hidden_states: torch.Tensor, + expert_ids: torch.Tensor, + weights: torch.Tensor, + qlen: int, + ) -> torch.device: """Copy inputs to CPU buffer, return input device.""" input_device = hidden_states.device buffer.input_cpu[:qlen].copy_(hidden_states.to(torch.bfloat16), non_blocking=True) @@ -510,11 +573,11 @@ def sync_backward(self) -> Tuple[torch.Tensor, torch.Tensor]: def submit_backward_repack(self): if not self._weights_loaded or self.moe is None: return - if hasattr(self.moe, 'submit_backward_repack'): + if hasattr(self.moe, "submit_backward_repack"): self.moe.submit_backward_repack() def wait_backward_repack(self): if not self._weights_loaded or self.moe is None: return - if hasattr(self.moe, 'wait_backward_repack'): + if hasattr(self.moe, "wait_backward_repack"): self.moe.wait_backward_repack() diff --git a/kt-kernel/python/sft/config.py b/kt-kernel/python/sft/config.py index 35af4d3d6..0172ad5dd 100644 --- a/kt-kernel/python/sft/config.py +++ b/kt-kernel/python/sft/config.py @@ -75,6 +75,10 @@ class KTConfig: kt_lora_rank: int | None = None kt_lora_alpha: float | None = None + # Training mode + kt_train_mode: str | None = None # "lora" | "full" | "hybrid" + kt_full_weight_grad: bool | None = None # auto-set True when train_mode in (full, hybrid) + # LoRA Experts (GPU-side extra experts) kt_use_lora_experts: bool | None = None kt_lora_expert_num: int | None = None @@ -132,6 +136,10 @@ def __post_init__(self): self.kt_lora_alpha = _env_float("ACCELERATE_KT_LORA_ALPHA", None) if self.kt_lora_alpha is None and self.kt_lora_rank is not None: self.kt_lora_alpha = float(self.kt_lora_rank * 2) + if self.kt_train_mode is None: + self.kt_train_mode = os.environ.get("ACCELERATE_KT_TRAIN_MODE", "lora") + if self.kt_full_weight_grad is None: + self.kt_full_weight_grad = self.kt_train_mode in ("full", "hybrid") if self.kt_model_max_length is None: self.kt_model_max_length = _env_int("ACCELERATE_KT_MODEL_MAX_LENGTH", None) if self.kt_skip_expert_loading is None: diff --git a/kt-kernel/python/sft/layer.py b/kt-kernel/python/sft/layer.py index e4cb2b657..fc12949d4 100644 --- a/kt-kernel/python/sft/layer.py +++ b/kt-kernel/python/sft/layer.py @@ -13,6 +13,7 @@ import logging import os +from contextlib import nullcontext from typing import Any import torch @@ -62,7 +63,29 @@ def __init__( # 1. gate/router FIRST - keep original attribute name for PEFT compatibility router_attr = moe_config.router_attr # "gate" for Qwen3/DeepSeek - setattr(self, router_attr, getattr(original_moe, router_attr, None)) + original_router = getattr(original_moe, router_attr, None) + self._original_router = None # Set when router is not nn.Linear (e.g. TopKRouter) + + if original_router is not None and isinstance(original_router, nn.Linear): + # transformers <=4.x / some models: gate is nn.Linear - register directly. + setattr(self, router_attr, original_router) + elif original_router is not None and hasattr(original_router, "weight") and isinstance( + getattr(original_router, "weight"), nn.Parameter + ): + # transformers v5+: gate is a TopKRouter with nn.Parameter weight. + # Wrap it in nn.Linear so PEFT can discover and inject LoRA. + # The nn.Linear shares the same weight tensor - LoRA applied to it + # is equivalent to LoRA on the original gate. + router_weight = original_router.weight + router_linear = nn.Linear( + router_weight.shape[1], router_weight.shape[0], bias=False, + ) + router_linear.weight = router_weight # share the same parameter + setattr(self, router_attr, router_linear) + # Keep the original router for forward (top-k selection logic) + self._original_router = original_router + else: + setattr(self, router_attr, original_router) self._router_attr = router_attr # 2. experts SECOND (this is what PEFT targets for LoRA) @@ -84,10 +107,13 @@ def __init__( self._peft_lora_modules: dict[int, dict[str, tuple[nn.Module, nn.Module]]] | None = None self._lora_pointers_dirty = False + # Full weight grad mode (set during wrapping or kt_adapt_peft_lora) + self._full_weight_grad = getattr(wrapper, "_full_weight_grad", False) if wrapper is not None else False + def _apply(self, fn, recurse=True): # Protect experts from device transfer (PEFT LoRA should stay on CPU for KT) saved_experts = None - experts_attr = getattr(self, '_experts_attr', None) + experts_attr = getattr(self, "_experts_attr", None) if experts_attr is not None and getattr(self, experts_attr, None) is not None: saved_experts = getattr(self, experts_attr) @@ -103,6 +129,7 @@ def _apply(self, fn, recurse=True): def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: import torch.distributed as dist + dist_on = dist.is_initialized() and dist.get_world_size() > 1 rank = dist.get_rank() if dist.is_initialized() else 0 @@ -112,11 +139,12 @@ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: topk_ids, topk_weights = self._compute_routing(hidden_states) train_lora = self._peft_lora_modules is not None and len(self._peft_lora_modules) > 0 + full_weight_grad = self._full_weight_grad save_for_backward = ( self.training and torch.is_grad_enabled() - and (hidden_states.requires_grad or topk_weights.requires_grad or train_lora) + and (hidden_states.requires_grad or topk_weights.requires_grad or train_lora or full_weight_grad) ) use_autograd_path = save_for_backward save_for_backward_submit = use_autograd_path @@ -127,6 +155,11 @@ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: self.update_lora_pointers() self._lora_pointers_dirty = False + # In full_weight_grad mode, sync base weights after optimizer step + if full_weight_grad and getattr(self.wrapper, "_base_weights_dirty", False): + self.wrapper.update_base_weights() + self.wrapper._base_weights_dirty = False + gpu_output, all_qlens = self._submit_and_compute_gpu( hidden_states, topk_ids, @@ -141,11 +174,15 @@ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: if train_lora and self._peft_lora_modules: for expert_loras in self._peft_lora_modules.values(): for lora_A, lora_B in expert_loras.values(): - if hasattr(lora_A, 'weight') and lora_A.weight.requires_grad: + if hasattr(lora_A, "weight") and lora_A.weight.requires_grad: lora_ref = lora_A.weight break if lora_ref.numel() > 0: break + elif full_weight_grad and self.wrapper is not None: + # In full mode, use base weight param as autograd sentinel + if self.wrapper.gate_proj_buf is not None: + lora_ref = self.wrapper.gate_proj_buf moe_output = KTMoEFunction.apply( hidden_states, @@ -159,6 +196,10 @@ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: save_for_backward, train_lora, all_qlens, + # Base weight params for full mode gradient flow + self.wrapper.gate_proj_buf if full_weight_grad and self.wrapper is not None else None, + self.wrapper.up_proj_buf if full_weight_grad and self.wrapper is not None else None, + self.wrapper.down_proj_buf if full_weight_grad and self.wrapper is not None else None, ) else: moe_output = self._sync_forward_output_no_autograd( @@ -194,13 +235,9 @@ def _sync_forward_output_no_autograd( else: all_qlens_list = [int(q) for q in all_qlens] if len(all_qlens_list) != world_size: - raise RuntimeError( - f"all_qlens length mismatch: got {len(all_qlens_list)}, expected {world_size}" - ) + raise RuntimeError(f"all_qlens length mismatch: got {len(all_qlens_list)}, expected {world_size}") if int(all_qlens_list[rank]) != qlen: - raise RuntimeError( - f"Rank {rank} qlen mismatch: local={qlen}, all_qlens[{rank}]={all_qlens_list[rank]}" - ) + raise RuntimeError(f"Rank {rank} qlen mismatch: local={qlen}, all_qlens[{rank}]={all_qlens_list[rank]}") total_qlen = sum(all_qlens_list) if rank == 0: @@ -234,12 +271,11 @@ def _sync_forward_output_no_autograd( return torch.empty(batch_size, seq_len, self.hidden_size, device=original_device, dtype=original_dtype) def _compute_routing(self, hidden_states: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: - # Run routing under no_grad to avoid creating autograd nodes whose - # SavedVariables become orphan holders inside gradient checkpoint. - # The gate is frozen during LoRA fine-tuning and the main gradient - # flows through KTMoEFunction.backward()'s grad_input, so the - # routing gradient contribution to hidden_states can be safely dropped. - with torch.no_grad(): + # In full_weight_grad mode, Router gradients should flow (no torch.no_grad). + # In LoRA mode, Router is frozen — wrap in no_grad to avoid orphan autograd nodes. + no_grad_ctx = torch.no_grad() if not self._full_weight_grad else nullcontext() + + with no_grad_ctx: router = getattr(self, self._router_attr) if self.router_type == "deepseek_gate": # DeepSeek V3's MoEGate has `assert not self.training` in its noaux_tc @@ -259,6 +295,60 @@ def _compute_routing(self, hidden_states: torch.Tensor) -> tuple[torch.Tensor, t topk_weights = topk_weights.to(torch.bfloat16) return topk_ids, topk_weights + # When _original_router is set, self.gate is an nn.Linear wrapper + # around the TopKRouter's weight. Use it (with PEFT LoRA if + # applied) for the linear projection, then replicate top-k logic. + if self._original_router is not None: + orig_router = self._original_router + router_logits = router(hidden_states.view(-1, self.hidden_size)) + if self.router_type == "glm4_moe_gate": + router_probs = torch.sigmoid(router_logits.float()) + correction_bias = getattr(orig_router, "e_score_correction_bias", None) + if correction_bias is None: + router_logits_for_choice = router_probs + else: + router_logits_for_choice = router_probs + correction_bias.to( + device=router_probs.device, + dtype=router_probs.dtype, + ) + n_group = getattr(orig_router, "n_group", 1) + topk_group = getattr(orig_router, "topk_group", n_group) + expert_num = self.moe_config.expert_num + group_scores = ( + router_logits_for_choice.view(-1, n_group, expert_num // n_group) + .topk(2, dim=-1)[0] + .sum(dim=-1) + ) + group_idx = torch.topk(group_scores, k=topk_group, dim=-1, sorted=False)[1] + group_mask = torch.zeros_like(group_scores) + group_mask.scatter_(1, group_idx, 1) + score_mask = ( + group_mask.unsqueeze(-1) + .expand(-1, n_group, expert_num // n_group) + .reshape(-1, expert_num) + ) + scores_for_choice = router_logits_for_choice.masked_fill(~score_mask.bool(), 0.0) + top_k = getattr(orig_router, "top_k", self.moe_config.num_experts_per_tok) + topk_ids = torch.topk(scores_for_choice, k=top_k, dim=-1, sorted=False)[1] + topk_weights = router_probs.gather(1, topk_ids) + if getattr(orig_router, "norm_topk_prob", True): + topk_weights = topk_weights / (topk_weights.sum(dim=-1, keepdim=True) + 1e-20) + topk_weights = topk_weights * getattr(orig_router, "routed_scaling_factor", 1.0) + if topk_weights.is_floating_point(): + topk_weights = topk_weights.to(torch.bfloat16) + return topk_ids, topk_weights + + router_probs = F.softmax(router_logits, dtype=torch.float, dim=-1) + top_k = getattr(orig_router, "top_k", self.moe_config.num_experts_per_tok) + norm_topk_prob = getattr(orig_router, "norm_topk_prob", True) + topk_weights, topk_ids = torch.topk(router_probs, top_k, dim=-1) + if norm_topk_prob: + topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True) + topk_weights = topk_weights.to(router_logits.dtype) + if topk_weights.is_floating_point(): + topk_weights = topk_weights.to(torch.bfloat16) + return topk_ids, topk_weights + router_output = router(hidden_states.view(-1, self.hidden_size)) # transformers v5 TopKRouter returns (router_logits, router_scores, router_indices) # directly — scores/indices are already topk-normalized. @@ -299,9 +389,7 @@ def _submit_and_compute_gpu( if dist_on: all_qlens = _all_gather_qlens(qlen, original_device, world_size) if int(all_qlens[rank]) != qlen: - raise RuntimeError( - f"Rank {rank} qlen mismatch: local={qlen}, all_qlens[{rank}]={all_qlens[rank]}" - ) + raise RuntimeError(f"Rank {rank} qlen mismatch: local={qlen}, all_qlens[{rank}]={all_qlens[rank]}") total_qlen = sum(all_qlens) hs_flat = hidden_states.view(qlen, self.hidden_size).contiguous() diff --git a/kt-kernel/python/sft/lora.py b/kt-kernel/python/sft/lora.py index 5a594ec8e..e9f8d6202 100644 --- a/kt-kernel/python/sft/lora.py +++ b/kt-kernel/python/sft/lora.py @@ -98,15 +98,10 @@ def _find_kt_wrappers(model: nn.Module): return wrappers -def get_kt_lora_params(model: nn.Module) -> list[nn.Parameter]: - """Get all MoE LoRA parameters from KT model. - - Returns PEFT LoRA parameters from expert modules and lora_experts parameters. - """ +def _collect_kt_lora_params(wrappers) -> list[nn.Parameter]: + """Collect LoRA-only trainable parameters from KT wrappers.""" params: list[nn.Parameter] = [] - wrappers = _find_kt_wrappers(model) - if wrappers: for wrapper in wrappers: # PEFT LoRA parameters (from _peft_lora_modules) @@ -114,9 +109,9 @@ def get_kt_lora_params(model: nn.Module) -> list[nn.Parameter]: if peft_lora_modules is not None: for expert_loras in peft_lora_modules.values(): for lora_A, lora_B in expert_loras.values(): - if hasattr(lora_A, 'weight') and lora_A.weight.requires_grad: + if hasattr(lora_A, "weight") and lora_A.weight.requires_grad: params.append(lora_A.weight) - if hasattr(lora_B, 'weight') and lora_B.weight.requires_grad: + if hasattr(lora_B, "weight") and lora_B.weight.requires_grad: params.append(lora_B.weight) # Fused expert LoRA parameters (KT-managed, not PEFT) fused_params = getattr(wrapper, "_fused_expert_lora_params", None) @@ -129,6 +124,61 @@ def get_kt_lora_params(model: nn.Module) -> list[nn.Parameter]: return params +def _collect_kt_full_weight_params(wrappers) -> list[nn.Parameter]: + """Collect optimizer-visible base expert parameters for full/hybrid KT SFT.""" + params: list[nn.Parameter] = [] + + if wrappers: + for wrapper in wrappers: + if getattr(wrapper, "_full_weight_grad", False) and wrapper.wrapper is not None: + if wrapper.wrapper.gate_proj_buf is not None: + params.append(wrapper.wrapper.gate_proj_buf) + if wrapper.wrapper.up_proj_buf is not None: + params.append(wrapper.wrapper.up_proj_buf) + if wrapper.wrapper.down_proj_buf is not None: + params.append(wrapper.wrapper.down_proj_buf) + + return params + + +def get_kt_lora_params(model: nn.Module) -> list[nn.Parameter]: + """Get KT parameters for legacy Trainer optimizer injection. + + Historically the patched Trainer calls this function after optimizer + creation. In full_weight_grad mode, returning only LoRA params silently + drops expert base weights from the optimizer, so this compatibility entry + point delegates to the full trainable collector when needed. + """ + wrappers = _find_kt_wrappers(model) + if not wrappers: + return [] + + if any(getattr(w, "_full_weight_grad", False) for w in wrappers): + return _collect_kt_full_weight_params(wrappers) + _collect_kt_lora_params(wrappers) + + return _collect_kt_lora_params(wrappers) + + +def get_kt_trainable_params(model: nn.Module) -> list[nn.Parameter]: + """Get all trainable parameters from KT model based on training mode. + + In full mode: returns base weight nn.Parameter buffers from wrappers. + In LoRA mode: returns LoRA parameters (same as get_kt_lora_params). + """ + wrappers = _find_kt_wrappers(model) + if not wrappers: + return [] + + # Check if any wrapper is in full_weight_grad mode + has_full_weight_grad = any(getattr(w, "_full_weight_grad", False) for w in wrappers) + + if has_full_weight_grad: + return _collect_kt_full_weight_params(wrappers) + _collect_kt_lora_params(wrappers) + else: + # LoRA mode: return LoRA parameters + return _collect_kt_lora_params(wrappers) + + # ============================================================================= # PEFT LoRA Adaptation # ============================================================================= @@ -175,8 +225,24 @@ def kt_adapt_peft_lora(model: nn.Module) -> None: # wrap as nn.Parameter for optimizer, and pre-assign .grad for C++ backward. if getattr(wrapper, "_fused_experts", False): lora_rank = getattr(wrapper, "_lora_rank", 1) + + # In full mode (lora_rank=0), skip LoRA buffer creation entirely. + # C++ kernel will not compute LoRA contributions when lora_rank=0. + if lora_rank == 0: + wrapper._fused_expert_lora_params = [] + wrapper._peft_lora_modules = None + logger.info( + f"[kt_adapt_peft_lora] Layer {layer_idx}: fused expert, " + f"full mode (lora_rank=0, no LoRA buffers)" + ) + adapted_count += 1 + continue + lora_buffers, lora_grad_buffers, lora_params = _create_fused_expert_lora_buffers( - wrapper, moe_config, lora_rank, torch.bfloat16, + wrapper, + moe_config, + lora_rank, + torch.bfloat16, ) if is_rank_0 and wrapper.wrapper is not None: @@ -197,6 +263,18 @@ def kt_adapt_peft_lora(model: nn.Module) -> None: if len(experts) == 0: continue + # In full mode (lora_rank=0), PEFT does not inject LoRA on experts. + # Skip LoRA detection and initialization entirely. + if getattr(wrapper, "_lora_rank", 1) == 0: + wrapper._peft_lora_modules = None + wrapper._fused_expert_lora_params = [] + logger.info( + f"[kt_adapt_peft_lora] Layer {layer_idx}: non-fused expert, " + f"full mode (lora_rank=0, no LoRA)" + ) + adapted_count += 1 + continue + # Collect references to PEFT LoRA modules for each expert # Structure: {expert_idx: {proj_name: (lora_A_module, lora_B_module)}} peft_lora_modules = {} @@ -228,7 +306,16 @@ def kt_adapt_peft_lora(model: nn.Module) -> None: # Store PEFT LoRA references on wrapper wrapper._peft_lora_modules = peft_lora_modules + # In full_weight_grad mode, PEFT LoRA is not injected by LlamaFactory, + # so no PEFT LoRA found is expected — skip the error. if not peft_lora_modules: + if getattr(wrapper, "_full_weight_grad", False): + logger.info( + f"[kt_adapt_peft_lora] Layer {layer_idx}: No PEFT LoRA found " + f"(full_weight_grad mode — expected, skipping)" + ) + adapted_count += 1 + continue raise RuntimeError( f"[kt_adapt_peft_lora] Layer {layer_idx}: No PEFT LoRA found on any expert. " f"Check that PEFT lora_target includes expert modules." @@ -510,9 +597,16 @@ def _replace_peft_weights_with_views( "[_replace_peft_weights_with_views] first param: " "id %s->%s (same=%s) data_ptr %s->%s buf_ptr=%s (match=%s) " "has_grad=%s requires_grad=%s shape=%s", - _old_id_a, _new_id_a, _old_id_a == _new_id_a, - _old_ptr_a, _new_ptr_a, _buf_ptr_a, _new_ptr_a == _buf_ptr_a, - _has_grad, lora_A.weight.requires_grad, tuple(lora_A.weight.shape), + _old_id_a, + _new_id_a, + _old_id_a == _new_id_a, + _old_ptr_a, + _new_ptr_a, + _buf_ptr_a, + _new_ptr_a == _buf_ptr_a, + _has_grad, + lora_A.weight.requires_grad, + tuple(lora_A.weight.shape), ) _first_logged = True _replaced += 1 @@ -526,12 +620,15 @@ def _replace_peft_weights_with_views( def update_kt_lora_pointers(model: nn.Module): - """Mark KT wrapper LoRA pointers as dirty after optimizer.step().""" + """Mark KT wrapper LoRA pointers and base weight pointers as dirty after optimizer.step().""" wrappers = _find_kt_wrappers(model) if wrappers: for wrapper in wrappers: wrapper._lora_pointers_dirty = True + # In full mode, base weights also need re-sync after optimizer step + if getattr(wrapper, "_full_weight_grad", False) and wrapper.wrapper is not None: + wrapper.wrapper._base_weights_dirty = True # ============================================================================= @@ -541,12 +638,10 @@ def update_kt_lora_pointers(model: nn.Module): def sync_kt_lora_gradients(model: nn.Module) -> None: """ - Synchronize KT-managed LoRA gradients across ranks. + Synchronize KT-managed gradients across ranks. - KT computes expert LoRA gradients only on rank 0 (gather/scatter path). This function broadcasts the - per-layer contiguous grad buffers from rank 0 to all ranks so that: - - gradient clipping sees identical grads on every rank - - optimizer.step() applies identical updates + In LoRA mode: synchronizes LoRA gradients only. + In full mode: synchronizes both base weight and LoRA gradients. """ import torch.distributed as dist @@ -557,17 +652,36 @@ def sync_kt_lora_gradients(model: nn.Module) -> None: if world_size <= 1: return - params = get_kt_lora_params(model) + # Sync base weight gradients in full mode + wrappers = _find_kt_wrappers(model) + if wrappers: + for wrapper in wrappers: + if not getattr(wrapper, "_full_weight_grad", False): + continue + if wrapper.wrapper is None: + continue + for grad_buf in ( + wrapper.wrapper.grad_gate_proj_buf, + wrapper.wrapper.grad_up_proj_buf, + wrapper.wrapper.grad_down_proj_buf, + ): + if grad_buf is not None: + grad_gpu = grad_buf.cuda() + dist.all_reduce(grad_gpu, op=dist.ReduceOp.SUM) + grad_gpu.div_(world_size) + grad_buf.copy_(grad_gpu.cpu()) + + # Sync LoRA gradients. Use the LoRA-only helper here because base gradients + # were synchronized above and get_kt_lora_params() is full-aware for legacy + # optimizer injection compatibility. + params = _collect_kt_lora_params(wrappers) if not params: return for param in params: if param.grad is not None: - # Move grad to the same device as the parameter for all-reduce - # Then move back to CPU original_device = param.grad.device if original_device.type == "cpu": - # All-reduce on CPU might be slow; consider using a GPU buffer grad_gpu = param.grad.cuda() dist.all_reduce(grad_gpu, op=dist.ReduceOp.SUM) grad_gpu.div_(world_size) diff --git a/kt-kernel/python/sft/weights.py b/kt-kernel/python/sft/weights.py index c15e22638..e535e1c48 100644 --- a/kt-kernel/python/sft/weights.py +++ b/kt-kernel/python/sft/weights.py @@ -108,9 +108,16 @@ def get_weight_tensor(mod): return gate_proj, up_proj, down_proj -def _clear_original_expert_weights(moe_module: nn.Module, moe_config: MOEArchConfig) -> None: +def _clear_original_expert_weights( + moe_module: nn.Module, moe_config: MOEArchConfig, full_weight_grad: bool = False +) -> None: """ Clear original expert weights to free memory after KT weights are loaded. + + In full_weight_grad mode, gate_proj_buf/up_proj_buf/down_proj_buf serve as + the authoritative copies for the optimizer. The original expert weights in + the model tree are redundant and cause double-counting in count_parameters(). + Clear them just like in LoRA mode. """ from .arch import detect_fused_experts @@ -127,10 +134,14 @@ def _clear_original_expert_weights(moe_module: nn.Module, moe_config: MOEArchCon original_dtype = param.dtype tiny_storage = torch.UntypedStorage(1, device="cpu") fake_tensor = torch.tensor([], dtype=original_dtype, device="cpu").set_( - tiny_storage, storage_offset=0, size=param.shape, + tiny_storage, + storage_offset=0, + size=param.shape, stride=[0] * len(param.shape), ) - experts._parameters[name] = nn.Parameter(fake_tensor, requires_grad=False) + placeholder = nn.Parameter(fake_tensor, requires_grad=False) + placeholder._kt_zero_storage = True # Mark for _setup_full_tuning / count_parameters to skip + experts._parameters[name] = placeholder return def _iter_weight_params(): @@ -141,7 +152,9 @@ def _iter_weight_params(): continue parametrizations = getattr(proj, "parametrizations", None) - parametrized_weight = getattr(parametrizations, "weight", None) if parametrizations is not None else None + parametrized_weight = ( + getattr(parametrizations, "weight", None) if parametrizations is not None else None + ) if parametrized_weight is not None: original = getattr(parametrized_weight, "original", None) if isinstance(original, torch.nn.Parameter): @@ -178,10 +191,13 @@ def _iter_weight_params(): # only used for shape/dtype discovery by PEFT. tiny_storage = torch.UntypedStorage(1, device="cpu") fake_tensor = torch.tensor([], dtype=original_dtype, device="cpu").set_( - tiny_storage, storage_offset=0, size=weight_param.shape, + tiny_storage, + storage_offset=0, + size=weight_param.shape, stride=[0] * len(weight_param.shape), ) new_param = nn.Parameter(fake_tensor, requires_grad=False) + new_param._kt_zero_storage = True # Mark for _setup_full_tuning / count_parameters to skip replaced_count += 1 # Avoid `KeyError: attribute 'weight' already exists` for parametrized modules @@ -201,9 +217,7 @@ def _iter_weight_params(): try: setattr(container, param_name, new_param) except Exception as exc: - logger.warning( - f"Failed to clear expert weight {type(proj).__name__}.{param_name}: {exc}" - ) + logger.warning(f"Failed to clear expert weight {type(proj).__name__}.{param_name}: {exc}") logger.info(f"Replaced {replaced_count} expert weight params") @@ -256,7 +270,9 @@ def _load_kt_weight_index(kt_weight_path: str) -> dict[str, str]: return index -def _dequant_fp8_experts(weights: list[torch.Tensor], scales: list[torch.Tensor | None], block_size: tuple[int, int]) -> torch.Tensor: +def _dequant_fp8_experts( + weights: list[torch.Tensor], scales: list[torch.Tensor | None], block_size: tuple[int, int] +) -> torch.Tensor: """Dequantize a list of FP8 expert weights and stack them (batched, vectorized). Args: @@ -468,9 +484,7 @@ def load_experts_from_kt_weight_path( f"Expected keys like 'blk.{layer_idx}.ffn_gate_exps.0.numa.0.weight'" ) - logger.info( - f"Loading INT8 weights for layer {layer_idx}: {num_experts} experts, {numa_count} NUMA partitions" - ) + logger.info(f"Loading INT8 weights for layer {layer_idx}: {num_experts} experts, {numa_count} NUMA partitions") gate_weights_list = [] gate_scales_list = [] diff --git a/kt-kernel/python/sft/wrapper.py b/kt-kernel/python/sft/wrapper.py index a53ea88af..bb5284353 100644 --- a/kt-kernel/python/sft/wrapper.py +++ b/kt-kernel/python/sft/wrapper.py @@ -95,9 +95,7 @@ def build_kt_device_map(config, kt_plugin, device: str = "cuda:0") -> dict[str, else: device_map[expert_key] = "cpu" - logger.info( - f"Built KT device_map: {num_gpu_experts} GPU experts, {num_experts - num_gpu_experts} CPU experts" - ) + logger.info(f"Built KT device_map: {num_gpu_experts} GPU experts, {num_experts - num_gpu_experts} CPU experts") return device_map @@ -163,8 +161,24 @@ def wrap_moe_layers_with_kt_wrapper(model: nn.Module, kt_plugin: Any) -> list[KT cfg = _get_kt_config(kt_plugin) # Read lora_rank/lora_alpha for C++ wrapper initialization (buffer allocation only) - lora_rank = getattr(cfg, "kt_lora_rank", 1) or 1 - lora_alpha = getattr(cfg, "kt_lora_alpha", 1.0) or 1.0 + # Use explicit None checks: lora_rank=0 is a valid value (full mode, no LoRA), + # but `or` pattern would treat 0 as falsy and replace it with 1. + _raw_rank = getattr(cfg, "kt_lora_rank", None) + lora_rank = _raw_rank if _raw_rank is not None else 1 + _raw_alpha = getattr(cfg, "kt_lora_alpha", None) + lora_alpha = _raw_alpha if _raw_alpha is not None else 1.0 + + # Read full_weight_grad mode + _raw_fwg = getattr(cfg, "kt_full_weight_grad", None) + full_weight_grad = _raw_fwg if _raw_fwg is not None else False + + # In full mode, lora_rank should be 0 (no LoRA, only base weight grad) + # If user explicitly set lora_rank > 0 in full mode (hybrid), keep it. + # Otherwise, auto-set lora_rank=0. + if full_weight_grad and lora_rank > 0: + _has_explicit_lora_rank = getattr(cfg, "kt_lora_rank", None) is not None + if not _has_explicit_lora_rank: + lora_rank = 0 # Read LoRA Experts configuration _raw_le = getattr(cfg, "kt_use_lora_experts", None) @@ -177,6 +191,8 @@ def wrap_moe_layers_with_kt_wrapper(model: nn.Module, kt_plugin: Any) -> list[KT f"LoRA Experts config: use_lora_experts={use_lora_experts}, " f"num={lora_expert_num}, intermediate_size={lora_expert_intermediate_size}" ) + if full_weight_grad: + logger.info(f"Full weight gradient mode enabled (lora_rank={lora_rank})") wrappers: list[KTMoELayerWrapper] = [] moe_layer_count = 0 @@ -225,7 +241,9 @@ def wrap_moe_layers_with_kt_wrapper(model: nn.Module, kt_plugin: Any) -> list[KT cfg.kt_sharded_metadata = sharded_metadata logger.info(f"Resolved {len(checkpoint_files)} checkpoint files from kt_expert_checkpoint_path") else: - logger.warning(f"Failed to resolve checkpoint files from kt_expert_checkpoint_path={kt_expert_checkpoint_path!r}") + logger.warning( + f"Failed to resolve checkpoint files from kt_expert_checkpoint_path={kt_expert_checkpoint_path!r}" + ) use_checkpoint_files = bool(checkpoint_files) and not use_kt_weight_path @@ -260,6 +278,7 @@ def wrap_moe_layers_with_kt_wrapper(model: nn.Module, kt_plugin: Any) -> list[KT ) import torch.distributed as _dist + _rank = _dist.get_rank() if _dist.is_initialized() else 0 model_container, layers = _get_model_container_and_layers(model, purpose="wrapping") @@ -329,6 +348,7 @@ def wrap_moe_layers_with_kt_wrapper(model: nn.Module, kt_plugin: Any) -> list[KT lora_rank=lora_rank, lora_alpha=lora_alpha, max_cache_depth=getattr(cfg, "kt_max_cache_depth", 2), + full_weight_grad=full_weight_grad, ) # Set share_backward_bb and share_cache_pool BEFORE load_weights (config is built during load) @@ -352,9 +372,18 @@ def wrap_moe_layers_with_kt_wrapper(model: nn.Module, kt_plugin: Any) -> list[KT physical_to_logical_map_cpu=physical_to_logical_map, ) - wrapper.gate_proj = None - wrapper.up_proj = None - wrapper.down_proj = None + # In full_weight_grad mode, keep weight references for nn.Parameter initialization + # and initialize the base weight buffers + if full_weight_grad: + wrapper.init_full_weight_grad_buffers( + gate_proj=wrapper.gate_proj if wrapper.gate_proj is not None else gate_proj, + up_proj=wrapper.up_proj if wrapper.up_proj is not None else up_proj, + down_proj=wrapper.down_proj if wrapper.down_proj is not None else down_proj, + ) + else: + wrapper.gate_proj = None + wrapper.up_proj = None + wrapper.down_proj = None # Create LoRA Experts if enabled lora_experts = None @@ -381,16 +410,18 @@ def wrap_moe_layers_with_kt_wrapper(model: nn.Module, kt_plugin: Any) -> list[KT setattr(layer, moe_config.moe_layer_attr, layer_wrapper) # Base weights have been copied into the C++ kernel's internal BufferB format. - # Do not hold a Python-side reference --- it wastes ~1 GB/layer. + # In full_weight_grad mode, the authoritative copies are gate_proj_buf etc. + # Always release local references to save ~1 GB/layer. del gate_proj, up_proj, down_proj wrappers.append(layer_wrapper) moe_layer_count += 1 - # Replace original expert weights with meta placeholders. + # Replace original expert weights with zero-storage placeholders. # Experts remain in the model tree (via wrapper.experts) so PEFT can discover them. # Rank 0 already copied weights to C++ kernel via load_weights_from_tensors. - _clear_original_expert_weights(moe_module, moe_config) + # gate_proj_buf serves as the authoritative copy in full_weight_grad mode. + _clear_original_expert_weights(moe_module, moe_config, full_weight_grad=full_weight_grad) logger.info(f"Wrapped {moe_layer_count} MoE layers with KTMoEWrapper") @@ -420,6 +451,17 @@ class should import it from the appropriate dataclasses module. from .config import KTConfig from accelerate.utils.dataclasses import KTransformersPlugin + # Map LlamaFactory finetuning_type to kt_train_mode + finetuning_type = getattr(finetuning_args, "finetuning_type", None) if finetuning_args else None + kt_train_mode_map = { + "full": "full", + "freeze": "hybrid", + "lora": "lora", + "galore": "full", + "badam": "full", + } + kt_train_mode = kt_train_mode_map.get(finetuning_type, None) if finetuning_type else None + kt_config = KTConfig( kt_backend=getattr(model_args, "kt_backend", None), kt_num_threads=getattr(model_args, "kt_num_threads", None), @@ -435,6 +477,7 @@ class should import it from the appropriate dataclasses module. kt_lora_rank=getattr(finetuning_args, "lora_rank", None) if finetuning_args else None, kt_lora_alpha=getattr(finetuning_args, "lora_alpha", None) if finetuning_args else None, kt_model_max_length=getattr(model_args, "model_max_length", None), + kt_train_mode=kt_train_mode, ) return KTransformersPlugin(enabled=True, kt_config=kt_config) @@ -509,7 +552,13 @@ def load_kt_model( **kwargs, ) -> nn.Module: """Load model with KTMoEWrapper backend.""" - from .arch import get_moe_arch_config, move_non_experts_to_gpu, get_expert_device, KTAMXNotAvailableError, KTAMXConfigError + from .arch import ( + get_moe_arch_config, + move_non_experts_to_gpu, + get_expert_device, + KTAMXNotAvailableError, + KTAMXConfigError, + ) if kt_plugin is None: if model_args is None: @@ -536,8 +585,11 @@ def load_kt_model( from transformers.integrations.kt import set_kt_config, unset_kt_config loading_kwargs = get_kt_loading_kwargs( - config, kt_plugin, torch_dtype=torch_dtype, - trust_remote_code=trust_remote_code, token=token, + config, + kt_plugin, + torch_dtype=torch_dtype, + trust_remote_code=trust_remote_code, + token=token, ) if model_args is not None: for key in ("cache_dir", "revision"): @@ -551,8 +603,10 @@ def load_kt_model( if getattr(cfg, "kt_skip_expert_loading", None) is None: checkpoint_files, sharded_metadata = _resolve_checkpoint_files( model_name_or_path=model_name_or_path, - cache_dir=cache_dir, revision=revision, - token=token, trust_remote_code=trust_remote_code, + cache_dir=cache_dir, + revision=revision, + token=token, + trust_remote_code=trust_remote_code, ) if checkpoint_files and all(f.endswith(".safetensors") for f in checkpoint_files): if getattr(cfg, "kt_weight_path", None) is None: From 2d81e8683a9da4b0394596adb820737e66465997 Mon Sep 17 00:00:00 2001 From: illu Date: Fri, 10 Jul 2026 11:04:29 +0000 Subject: [PATCH 02/20] [fix](kt-kernel): fix Full FT TP base weight gradients --- kt-kernel/AGENTS.md | 190 ++++++++++++++++++++++++++++ kt-kernel/operators/amx/sft_moe.hpp | 36 ++++-- kt-kernel/operators/moe-sft-tp.hpp | 44 ++++++- 3 files changed, 255 insertions(+), 15 deletions(-) create mode 100644 kt-kernel/AGENTS.md diff --git a/kt-kernel/AGENTS.md b/kt-kernel/AGENTS.md new file mode 100644 index 000000000..b31484dcb --- /dev/null +++ b/kt-kernel/AGENTS.md @@ -0,0 +1,190 @@ +# AGENTS.md: kt-kernel Full FT 开发指南 + +本文件作用于 `kt-kernel/` 及其全部子目录,用于指导 Agent 开发和验证 MoE +Full Fine-Tuning(Full FT)功能。重点是专家基座权重的梯度、optimizer 更新 +和 AMX 权重重新量化,不讨论普通 GPU 参数训练。 + +当前开发快照: + +- 仓库:`Illumination111/ktransformers-fullFT_development` +- Full FT 基线提交:`e99b5e1`,基于上游 `8e46e58` +- Full FT 来源:poryfly 的 `239bac5`、`25fa2fd` +- 当前状态:TP 子核配置、梯度分片/stride、临时区容量和逐步清零已修复;等待修复后训练验证 + +## 1. 开发前先确认 + +1. 在 `ktransformers` 中运行 `git status --short --branch`,不要覆盖已有改动。 +2. 使用符号名定位代码,不依赖本文行号。 +3. 先确认目标是 `full`、`hybrid` 还是 `lora`: + - `full`:训练 expert 基座权重,通常 `lora_rank=0`。 + - `hybrid`:同时训练基座权重和 LoRA。 + - `lora`:不得计算或更新基座权重。 +4. 修改前记录 TP 数量、NUMA 数量、`E/H/F` 和权重 dtype。 + +## 2. Full FT 数据契约 + +完整链路: + +```text +KTConfig + -> wrapper 创建 CPU BF16 nn.Parameter + -> forward 将 Parameter 挂入 KTMoEFunction + -> C++ backward 写 grad_*_proj_buf + -> autograd 返回基座梯度 + -> optimizer 更新 *_proj_buf + -> update_base_weights 重新量化到 AMX BufferB +``` + +权威数据必须唯一: + +- `gate/up/down_proj_buf`:optimizer 可见的 CPU BF16 基座权重。 +- `grad_gate/up/down_proj_buf`:C++ backward 的输出。 +- AMX 量化权重:forward 使用的派生副本,optimizer 后必须重建。 +- HF expert 原权重只是占位,不得作为 Full FT 的更新依据。 + +## 3. 文件职责 + +路径均相对本文件所在的 `kt-kernel/`。 + +| 路径 | 职责 | +|---|---| +| `python/sft/config.py` | 将 `full/hybrid` 映射为 `kt_full_weight_grad=True` | +| `python/sft/base.py` | 创建 `*_proj_buf`、`grad_*_proj_buf` | +| `python/sft/wrapper.py` | 包装 MoE、确定权威权重 | +| `python/sft/layer.py` | 将三个基座 Parameter 传入 autograd;dirty 时 requant | +| `python/sft/autograd.py` | 返回 C++ 写入的基座梯度 | +| `python/sft/lora.py` | optimizer 参数注入、分布式同步、dirty 标记 | +| `python/sft/amx.py` | 设置 C++ config、传梯度指针、触发重新量化 | +| `operators/common.hpp` | `MOESFTConfig.full_weight_grad` | +| `operators/moe-tp.hpp` | 通用 TP 切分;当前会把 SFT config 切成基类 | +| `operators/moe-sft-tp.hpp` | SFT TP 调度和全局梯度 buffer 分片 | +| `operators/amx/sft_moe.hpp` | NUMA 子核 backward 和基座梯度计算 | +| `ext_bindings.cpp` | Python/C++ 参数和 task 绑定 | + +## 4. TP 张量布局 + +定义: + +- `E`:expert 数量。 +- `H`:hidden size。 +- `F`:完整 `intermediate_size`。 +- `I`:当前 TP 子核的本地 `intermediate_size`。 +- `tp_offset`:当前 TP 在完整 `F` 维上的起点。 + +完整梯度布局: + +```text +gate/up: [E, F, H] +down: [E, H, F] +``` + +TP wrapper 传给子核的起始指针: + +```text +gate/up = base + tp_offset * H +down = base + tp_offset +``` + +子核循环和 FP32 累加使用本地 `I`,但写回完整张量时必须使用 `F` +作为 expert/row stride: + +```text +gate/up: expert * F * H + local_i * H + h +down: expert * H * F + h * F + local_i +``` + +所有 TP 必须写入互不重叠的区域;对空指针不得做偏移运算。 + +## 5. 历史根因与失败基线 + +`full_weight_grad=True` 已正确到达顶层 `TP_MOE_SFT`,但 +`TP_MOE` 创建子核时执行: + +```cpp +GeneralMOEConfig tp_config = config; +``` + +这会丢失 `MOESFTConfig.full_weight_grad`。子核随后通过 +`MOESFTConfig(GeneralMOEConfig)` 重建配置,字段回到默认 `false`, +因此 `backward_base_weight_grad` 被门控跳过。 + +只恢复该布尔值仍不够:当前 TP backward 把相同的完整梯度首地址传给 +所有子核,子核又按本地 `I` 写回,会导致覆盖和错误布局。 + +修复前的 `20260710_103230_1gpu_AMX_BF16_FULLFT_EXPERTONLY` run 已冻结 +其他权重并只向 optimizer 注入 144 个 expert 基座参数,但 15 步 +`grad_norm=0`,step 0 到 step 15 的 probe 为 `changed=0/12`、 +`max_abs_delta=0`。720 次 requant 不能证明权重更新。 + +## 6. 当前两文件实现 + +生产代码限制在两个文件,未修改 Python 接口。 + +### 6.1 `operators/amx/sft_moe.hpp` + +1. `set_full_weight_grad(bool)` 更新子核的 + `sft_config_.full_weight_grad`。 +2. `full_intermediate_size` 已传给 `backward_base_weight_grad`。 +3. 计算循环使用本地 `I`,写回 stride 使用完整 `F`。 +4. 门控要求开关为 true 且三个梯度指针均非空。 +5. `forward_pool_` 保证至少 `3 * I * H * sizeof(float)` 字节。 +6. `lora_rank=0` 时 scaling 明确为 0,避免纯 Full FT 产生 Inf。 + +### 6.2 `operators/moe-sft-tp.hpp` + +1. 子核创建后逐个调用 `set_full_weight_grad`。 +2. 传播逻辑位于 `if constexpr (!kSkipLoRA)` 外;纯 Full FT 的 + `lora_rank=0` 也必须执行。 +3. backward dispatch 前按第 4 节公式生成每个 TP 的三个梯度指针。 +4. 完整 `F` 传给子核作为全局 stride,并检查所有 TP slice 完整覆盖 `F`。 +5. TP dispatch 前统一并行清零三组完整梯度,避免跨 step 残留和 TP 清零竞争。 + +暂不修改通用 `operators/moe-tp.hpp`。保留派生 config 的泛型重构影响 +所有 MoE backend,应作为后续独立改动。 + +## 7. 验证顺序 + +当前已通过 `clang-format --dry-run --Werror` 和 AMX/CUDA Release +`build_ext --inplace`。修复后的 expert-only 短训练仍是必需验收项。 + +1. **静态检查**:确认 setter 不在 LoRA 条件分支内;确认所有指针先判空再偏移。 +2. **构建**:重新编译并安装 `kt_kernel_ext`。 +3. **参考梯度测试**: + - 分别覆盖 TP=1 和 TP=2。 + - 使用 `expert_id > 0`,避免错误 expert stride 被 expert 0 掩盖。 + - 与 PyTorch outer-product 参考值比较 gate/up/down 三组梯度。 + - TP=2 时分别检查两个 `F` 分片,确认无覆盖。 +4. **跨 step 测试**:连续两步激活不同 expert,确认未激活 expert 不保留旧梯度。 +5. **模式回归**:`full_weight_grad=false` 不写基座梯度;纯 LoRA 结果不变。 +6. **短训练**: + - step 内三个 `grad_*_proj_buf.abs().max() > 0`。 + - optimizer 后至少一个 `*_proj_buf` 发生有限变化。 + - requant 后下一次 forward 使用更新后的权重。 + +不要用 loss 下降证明 Full FT 生效;attention、router 或其他非 expert 参数也能 +让 loss 下降。 + +## 8. 完成标准 + +- C++ 子核实际收到 `full_weight_grad=true`。 +- 三组基座梯度非零、有限,并与参考实现一致。 +- TP 分片覆盖完整 `F`,无重叠、越界或错误 stride。 +- optimizer 能找到三个基座 Parameter,并在 step 后更新它们。 +- inactive expert 无跨 step 残留梯度。 +- `full`、`hybrid`、`lora` 三种模式行为符合各自契约。 +- 纯 LoRA 和现有 AMX MoE 测试无回归。 + +## 9. 调试资料 + +```text +FFTtest/Qwen3-30B-A3B/test_log/20260710_103230_1gpu_AMX_BF16_FULLFT_EXPERTONLY/ + summary.md + expert_weight_change_check.{txt,json} + phase4/expert_buf_probe.json + phase4/train.log + +FFTtest/Qwen3-30B-A3B/expert_buf_probe.py +``` + +该 run 的直接证据是 `changed=0/12`、`max_abs_delta=0`;训练 loss +仍下降不能否定 expert 基座梯度失效。 diff --git a/kt-kernel/operators/amx/sft_moe.hpp b/kt-kernel/operators/amx/sft_moe.hpp index b516bcf75..dd8316151 100644 --- a/kt-kernel/operators/amx/sft_moe.hpp +++ b/kt-kernel/operators/amx/sft_moe.hpp @@ -712,6 +712,13 @@ class AMX_SFT_MOE_TP : public BaseMOE { work_required += round_up(lora_bc_out_pool_bytes_, kAmxAlignment); work_required += round_up(lora_intermediate_bf16_pool_bytes_, kAmxAlignment); + // Base-weight gradients reuse the forward pool after LoRA backward is complete. + if (sft_config_.full_weight_grad) { + const size_t base_grad_accumulator_bytes = + 3 * (size_t)config_.intermediate_size * config_.hidden_size * sizeof(float); + work_required = std::max(work_required, base_grad_accumulator_bytes); + } + alloc_or_resize_forward_pool(work_required); SFT_POOL_LOG("fwd_work", config_.layer_idx, tp_part_idx, 0, cache_stack_top_, forward_pool_bytes_, @@ -897,9 +904,11 @@ class AMX_SFT_MOE_TP : public BaseMOE { */ void set_lora_params(int rank, float alpha) { lora_rank_ = rank; - lora_scaling_ = alpha / rank; + lora_scaling_ = rank > 0 ? alpha / rank : 0.0f; } + void set_full_weight_grad(bool enabled) { sft_config_.full_weight_grad = enabled; } + /** * @brief SFT Forward pass with optional caching for backward. * @@ -1883,7 +1892,8 @@ class AMX_SFT_MOE_TP : public BaseMOE { // Step 5: Base weight gradient accumulation (full weight grad mode) // ===================================================================== if (sft_config_.full_weight_grad && grad_gate_proj && grad_up_proj && grad_down_proj) { - backward_base_weight_grad(cache, grad_output, grad_gate_proj, grad_up_proj, grad_down_proj); + backward_base_weight_grad(cache, grad_output, full_intermediate_size, grad_gate_proj, grad_up_proj, + grad_down_proj); } // \u2605 Cache pool is NOT freed here \u2014 kept for reuse across steps. @@ -1905,16 +1915,20 @@ class AMX_SFT_MOE_TP : public BaseMOE { * * Uses FP32 accumulator for precision, writes BF16 output. */ - void backward_base_weight_grad(const ForwardCache& cache, const void* grad_output, void* grad_gate_proj, - void* grad_up_proj, void* grad_down_proj) { + void backward_base_weight_grad(const ForwardCache& cache, const void* grad_output, int full_intermediate_size, + void* grad_gate_proj, void* grad_up_proj, void* grad_down_proj) { const int H = config_.hidden_size; const int I = config_.intermediate_size; - const int E = config_.expert_num; + const int F = full_intermediate_size; int activated_expert = cache.activated_expert_cache; - auto* ggp = static_cast(grad_gate_proj); // [E, I, H] - auto* gup_ptr = static_cast(grad_up_proj); // [E, I, H] - auto* gdp = static_cast(grad_down_proj); // [E, H, I] + if (F < I) { + throw std::runtime_error("full_intermediate_size must be at least the TP-local intermediate_size"); + } + + auto* ggp = static_cast(grad_gate_proj); // TP slice of [E, F, H] + auto* gup_ptr = static_cast(grad_up_proj); // TP slice of [E, F, H] + auto* gdp = static_cast(grad_down_proj); // TP slice of [E, H, F] auto* grad_out_bf16 = static_cast(grad_output); for (int task_id = 0; task_id < activated_expert; task_id++) { @@ -1969,13 +1983,13 @@ class AMX_SFT_MOE_TP : public BaseMOE { // Convert FP32 accumulators to BF16 and store for (int i = 0; i < I; i++) { for (int h = 0; h < H; h++) { - ggp[(size_t)expert_idx * I * H + (size_t)i * H + h] = GGML_FP32_TO_BF16(acc_gate[i * H + h]); - gup_ptr[(size_t)expert_idx * I * H + (size_t)i * H + h] = GGML_FP32_TO_BF16(acc_up[i * H + h]); + ggp[(size_t)expert_idx * F * H + (size_t)i * H + h] = GGML_FP32_TO_BF16(acc_gate[i * H + h]); + gup_ptr[(size_t)expert_idx * F * H + (size_t)i * H + h] = GGML_FP32_TO_BF16(acc_up[i * H + h]); } } for (int h = 0; h < H; h++) { for (int i = 0; i < I; i++) { - gdp[(size_t)expert_idx * H * I + (size_t)h * I + i] = GGML_FP32_TO_BF16(acc_down[h * I + i]); + gdp[(size_t)expert_idx * H * F + (size_t)h * F + i] = GGML_FP32_TO_BF16(acc_down[h * I + i]); } } } diff --git a/kt-kernel/operators/moe-sft-tp.hpp b/kt-kernel/operators/moe-sft-tp.hpp index c8a0e0986..918aaf74d 100644 --- a/kt-kernel/operators/moe-sft-tp.hpp +++ b/kt-kernel/operators/moe-sft-tp.hpp @@ -229,6 +229,11 @@ class TP_MOE_SFT : public TP_MOE { part_grad_input_.assign(tp_count, nullptr); part_grad_weights_.assign(tp_count, nullptr); + // TP_MOE stores GeneralMOEConfig and slices off SFT-only fields. + for (int i = 0; i < tp_count; i++) { + tps[i]->set_full_weight_grad(config.full_weight_grad); + } + if constexpr (!kSkipLoRA) { // Bug #16 fix: TP_MOE base class uses GeneralMOEConfig (object slicing) which loses // LoRA pointers. We need to propagate LoRA pointers to all NUMA node instances. @@ -562,6 +567,8 @@ class TP_MOE_SFT : public TP_MOE { int k = sft_config.num_experts_per_tok; const bool need_grad_weights = (grad_weights != nullptr); + const bool need_base_weight_grad = sft_config.full_weight_grad && grad_gate_proj != nullptr && + grad_up_proj != nullptr && grad_down_proj != nullptr; // SkipLoRA: zero out lora_rank to skip all LoRA buffer allocations if constexpr (kSkipLoRA) lora_rank = 0; @@ -628,8 +635,7 @@ class TP_MOE_SFT : public TP_MOE { clear_bytes[i] = offset; } - // Parallel memset: zero only per-TP sparse partials and per-TP grad_input/grad_weights partials. - // The caller is responsible for passing zero-initialized final grad tensors. + // Parallel memset for per-TP partials and final base-weight gradients. struct ClearSeg { uint8_t* ptr; size_t len; @@ -650,6 +656,20 @@ class TP_MOE_SFT : public TP_MOE { } } + if (need_base_weight_grad) { + const size_t base_grad_bytes = (size_t)expert_num * full_intermediate_size * hidden_size * sizeof(ggml_bf16_t); + auto append_clear_segments = [&](void* ptr) { + auto* base = static_cast(ptr); + for (size_t off = 0; off < base_grad_bytes; off += kChunkBytes) { + size_t len = std::min(kChunkBytes, base_grad_bytes - off); + clear_segs.push_back(ClearSeg{base + off, len}); + } + }; + append_clear_segments(grad_gate_proj); + append_clear_segments(grad_up_proj); + append_clear_segments(grad_down_proj); + } + pool->do_work_stealing_job((int)clear_segs.size(), nullptr, [&](int seg_idx) { const auto& seg = clear_segs[(size_t)seg_idx]; @@ -665,6 +685,9 @@ class TP_MOE_SFT : public TP_MOE { std::vector tp_fp32_down_b(tp_count); std::vector tp_fp32_gate_a(tp_count); std::vector tp_fp32_up_a(tp_count); + std::vector tp_grad_gate_proj(tp_count, nullptr); + std::vector tp_grad_up_proj(tp_count, nullptr); + std::vector tp_grad_down_proj(tp_count, nullptr); if constexpr (!kSkipLoRA) { int tp_offset = 0; @@ -683,6 +706,19 @@ class TP_MOE_SFT : public TP_MOE { } } + int tp_offset = 0; + for (int i = 0; i < tp_count; i++) { + if (need_base_weight_grad) { + tp_grad_gate_proj[i] = static_cast(grad_gate_proj) + (size_t)tp_offset * hidden_size; + tp_grad_up_proj[i] = static_cast(grad_up_proj) + (size_t)tp_offset * hidden_size; + tp_grad_down_proj[i] = static_cast(grad_down_proj) + tp_offset; + } + tp_offset += tp_configs[i].intermediate_size; + } + if (tp_offset != full_intermediate_size) { + throw std::runtime_error("TP intermediate_size slices do not cover the full intermediate_size"); + } + // Run backward on each NUMA node pool->dispense_backend()->do_numa_job([&](int numa_id) { tps[numa_id]->backward(grad_output, part_grad_input_[numa_id], @@ -694,8 +730,8 @@ class TP_MOE_SFT : public TP_MOE { tp_down_a_ptr[numa_id], /* copy-type: direct write */ nullptr, /* grad_down_lora_b — unused, FP32 path below */ part_grad_weights_[numa_id], full_intermediate_size, tp_fp32_down_b[numa_id], - tp_fp32_gate_a[numa_id], tp_fp32_up_a[numa_id], grad_gate_proj, grad_up_proj, - grad_down_proj); + tp_fp32_gate_a[numa_id], tp_fp32_up_a[numa_id], tp_grad_gate_proj[numa_id], + tp_grad_up_proj[numa_id], tp_grad_down_proj[numa_id]); }); // // Collect per-thread timing from all NUMA subpools From 20f645c85c05b4cef31adc6070ac788944e5b567 Mon Sep 17 00:00:00 2001 From: illu Date: Mon, 13 Jul 2026 03:14:59 +0000 Subject: [PATCH 03/20] [fix]: bug fix of 2d81e86 --- kt-kernel/AGENTS.md | 57 +++++++++++++++++++++++++++-- kt-kernel/operators/amx/sft_moe.hpp | 15 +++----- 2 files changed, 59 insertions(+), 13 deletions(-) diff --git a/kt-kernel/AGENTS.md b/kt-kernel/AGENTS.md index b31484dcb..b32370906 100644 --- a/kt-kernel/AGENTS.md +++ b/kt-kernel/AGENTS.md @@ -9,7 +9,8 @@ Full Fine-Tuning(Full FT)功能。重点是专家基座权重的梯度、opt - 仓库:`Illumination111/ktransformers-fullFT_development` - Full FT 基线提交:`e99b5e1`,基于上游 `8e46e58` - Full FT 来源:poryfly 的 `239bac5`、`25fa2fd` -- 当前状态:TP 子核配置、梯度分片/stride、临时区容量和逐步清零已修复;等待修复后训练验证 +- 当前状态:TP 子核配置、梯度分片/stride、临时区容量和逐步清零已修复;GDB 已将 + 首步 SIGSEGV 定位为 base-weight backward 错用路由表,现改用 expert-major packed buffer ## 1. 开发前先确认 @@ -116,6 +117,24 @@ GeneralMOEConfig tp_config = config; `grad_norm=0`,step 0 到 step 15 的 probe 为 `changed=0/12`、 `max_abs_delta=0`。720 次 requant 不能证明权重更新。 +修复后的 `20260710_192456_1gpu_AMX_BF16_FULLFT_EXPERTONLY` run 未进入 +第一个训练 step。`phase4/train.log` 记录训练子进程 PID 3296907 在 +`2026-07-10 19:27:33 +08:00` 以 `exitcode=-11` 退出,并明确报告 +`Signal 11 (SIGSEGV)`;外层 `accelerate` 的 `exit_code.txt=1` 只是 launcher +退出码。`phase4/log_analysis.txt` 正确识别到崩溃,但旧 `summary.md` 的 +“未检测到崩溃”是汇总脚本误判,不得作为反证。该 run 没有 C++ backtrace、 +源码行号或 core,因而“C++ 梯度索引越界”目前只是诊断假设,并非已确认根因。 + +带符号 GDB 的 `20260713_105328_1gpu_AMX_BF16_FULLFT_EXPERTONLY` run +确认根因位于 `backward_base_weight_grad`:`m_local_pos_cache` 的真实布局是 +`[token_idx][route_slot]`,内层长度为 `k=8`,旧代码却按 +`m_local_pos_cache[expert_idx][t]` 访问。首个 expert 执行到 `t=8` 时越界, +读出垃圾 `tok_pos=81349952`,最终在读取 `input_row[0]` 时于 +`operators/amx/sft_moe.hpp:1968` 触发 SIGSEGV。两个 NUMA 子核均进入同一错误路径。 +最小修复不再反查路由表,而是直接使用 backward 已恢复/生成的 +`m_local_input_ptr_[expert_idx]` 和 `grad_output_bf16_ptr_[expert_idx]`;后者已包含 +router weight,因而同时保证 down projection 基座梯度的权重语义正确。 + ## 6. 当前两文件实现 生产代码限制在两个文件,未修改 Python 接口。 @@ -129,6 +148,8 @@ GeneralMOEConfig tp_config = config; 4. 门控要求开关为 true 且三个梯度指针均非空。 5. `forward_pool_` 保证至少 `3 * I * H * sizeof(float)` 字节。 6. `lora_rank=0` 时 scaling 明确为 0,避免纯 Full FT 产生 Inf。 +7. 基座梯度读取 expert-major packed input/grad-output,不把 + `m_local_pos_cache[token][route]` 误当作 expert token 列表。 ### 6.2 `operators/moe-sft-tp.hpp` @@ -145,7 +166,23 @@ GeneralMOEConfig tp_config = config; ## 7. 验证顺序 当前已通过 `clang-format --dry-run --Werror` 和 AMX/CUDA Release -`build_ext --inplace`。修复后的 expert-only 短训练仍是必需验收项。 +`build_ext --inplace`。但当前 Kllama 环境中的 `kt_kernel_ext` 已 stripped, +普通 GDB 只能得到地址或有限符号;定位 SIGSEGV 前应使用相同 Python 环境重新构建 +并安装带符号的 `RelWithDebInfo` 版本: + +```bash +cd /mnt/data2/wbw/ktransformers +CPUINFER_BUILD_TYPE=RelWithDebInfo \ + /mnt/data2/wbw/conda/envs/Kllama/bin/python3.12 \ + kt-kernel/setup.py build_ext --inplace +CPUINFER_BUILD_TYPE=RelWithDebInfo \ + /mnt/data2/wbw/conda/envs/Kllama/bin/python3.12 -m pip install \ + --no-build-isolation --no-deps --force-reinstall ./kt-kernel +``` + +确认测试实际 import 的 `.so` 含 `.debug_info`/`.debug_line`。然后使用 expert-only +runner 的 `--gdb`;只有磁盘空间足够时才使用 `--gdb-core`。修复后的 expert-only +短训练仍是必需验收项。 1. **静态检查**:确认 setter 不在 LoRA 条件分支内;确认所有指针先判空再偏移。 2. **构建**:重新编译并安装 `kt_kernel_ext`。 @@ -183,8 +220,22 @@ FFTtest/Qwen3-30B-A3B/test_log/20260710_103230_1gpu_AMX_BF16_FULLFT_EXPERTONLY/ phase4/expert_buf_probe.json phase4/train.log +FFTtest/Qwen3-30B-A3B/test_log/20260710_192456_1gpu_AMX_BF16_FULLFT_EXPERTONLY/ + phase4/train.log # exitcode=-11,Signal 11 (SIGSEGV) + phase4/log_analysis.txt # 正确检测到 SIGSEGV + phase4/exit_code.txt # 外层 launcher 退出码 1 + summary.md # P5 旧结论错误,不得作为崩溃判据 + +FFTtest/Qwen3-30B-A3B/test_log/20260713_105328_1gpu_AMX_BF16_FULLFT_EXPERTONLY/ + phase4/gdb_sigsegv.log # t=8 越界、非法 tok_pos 和双 NUMA 原生栈 + phase4/train.log + summary.md + FFTtest/Qwen3-30B-A3B/expert_buf_probe.py +FFTtest/Qwen3-30B-A3B/gdb_sigsegv.gdb +FFTtest/Qwen3-30B-A3B/run_full_ft_test_1gpu_bf16_frozen.sh --gdb ``` 该 run 的直接证据是 `changed=0/12`、`max_abs_delta=0`;训练 loss -仍下降不能否定 expert 基座梯度失效。 +仍下降不能否定 expert 基座梯度失效。后续 run 的直接崩溃证据必须以原始 +`phase4/train.log` 和 `phase4/gdb_sigsegv.log` 为准,不能只依赖生成的 summary。 diff --git a/kt-kernel/operators/amx/sft_moe.hpp b/kt-kernel/operators/amx/sft_moe.hpp index dd8316151..efb7d6f54 100644 --- a/kt-kernel/operators/amx/sft_moe.hpp +++ b/kt-kernel/operators/amx/sft_moe.hpp @@ -1892,8 +1892,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { // Step 5: Base weight gradient accumulation (full weight grad mode) // ===================================================================== if (sft_config_.full_weight_grad && grad_gate_proj && grad_up_proj && grad_down_proj) { - backward_base_weight_grad(cache, grad_output, full_intermediate_size, grad_gate_proj, grad_up_proj, - grad_down_proj); + backward_base_weight_grad(cache, full_intermediate_size, grad_gate_proj, grad_up_proj, grad_down_proj); } // \u2605 Cache pool is NOT freed here \u2014 kept for reuse across steps. @@ -1915,8 +1914,8 @@ class AMX_SFT_MOE_TP : public BaseMOE { * * Uses FP32 accumulator for precision, writes BF16 output. */ - void backward_base_weight_grad(const ForwardCache& cache, const void* grad_output, int full_intermediate_size, - void* grad_gate_proj, void* grad_up_proj, void* grad_down_proj) { + void backward_base_weight_grad(const ForwardCache& cache, int full_intermediate_size, void* grad_gate_proj, + void* grad_up_proj, void* grad_down_proj) { const int H = config_.hidden_size; const int I = config_.intermediate_size; const int F = full_intermediate_size; @@ -1929,7 +1928,6 @@ class AMX_SFT_MOE_TP : public BaseMOE { auto* ggp = static_cast(grad_gate_proj); // TP slice of [E, F, H] auto* gup_ptr = static_cast(grad_up_proj); // TP slice of [E, F, H] auto* gdp = static_cast(grad_down_proj); // TP slice of [E, H, F] - auto* grad_out_bf16 = static_cast(grad_output); for (int task_id = 0; task_id < activated_expert; task_id++) { int expert_idx = cache.m_expert_id_map_cache[task_id]; @@ -1941,8 +1939,6 @@ class AMX_SFT_MOE_TP : public BaseMOE { pos_start += cache.m_local_num_cache[cache.m_expert_id_map_cache[prev_id]]; } - const auto& local_pos = cache.m_local_pos_cache[expert_idx]; - // Allocate FP32 accumulators from forward pool (safe during backward) float* acc_gate = static_cast(forward_pool_); // [I, H] float* acc_up = acc_gate + (size_t)I * H; // [I, H] @@ -1953,12 +1949,11 @@ class AMX_SFT_MOE_TP : public BaseMOE { std::memset(acc_down, 0, (size_t)H * I * sizeof(float)); for (int t = 0; t < m; t++) { - int tok_pos = local_pos[t]; - const ggml_bf16_t* input_row = cache.input_cache + (size_t)tok_pos * H; + const ggml_bf16_t* input_row = m_local_input_ptr_[expert_idx] + (size_t)t * H; const ggml_bf16_t* gate_grad_row = grad_gate_output_ + (size_t)(pos_start + t) * I; const ggml_bf16_t* up_grad_row = grad_up_output_ + (size_t)(pos_start + t) * I; const ggml_bf16_t* inter_row = cache.intermediate_cache + (size_t)(pos_start + t) * I; - const ggml_bf16_t* grad_out_row = grad_out_bf16 + (size_t)tok_pos * H; + const ggml_bf16_t* grad_out_row = grad_output_bf16_ptr_[expert_idx] + (size_t)t * H; // gate_proj grad: [I, H] += grad_gate_out[t]^T @ input[t] for (int i = 0; i < I; i++) { From 5ce0767d1920db671f60dc1bd78243de588fb2c3 Mon Sep 17 00:00:00 2001 From: illu Date: Mon, 13 Jul 2026 05:35:27 +0000 Subject: [PATCH 04/20] [fix](kt-kernel): fix AMX BF16 full-weight gradients --- kt-kernel/AGENTS.md | 16 ++- kt-kernel/operators/amx/sft_moe.hpp | 193 ++++++++++++++++++++++++++-- 2 files changed, 198 insertions(+), 11 deletions(-) diff --git a/kt-kernel/AGENTS.md b/kt-kernel/AGENTS.md index b32370906..2648731b7 100644 --- a/kt-kernel/AGENTS.md +++ b/kt-kernel/AGENTS.md @@ -10,7 +10,9 @@ Full Fine-Tuning(Full FT)功能。重点是专家基座权重的梯度、opt - Full FT 基线提交:`e99b5e1`,基于上游 `8e46e58` - Full FT 来源:poryfly 的 `239bac5`、`25fa2fd` - 当前状态:TP 子核配置、梯度分片/stride、临时区容量和逐步清零已修复;GDB 已将 - 首步 SIGSEGV 定位为 base-weight backward 错用路由表,现改用 expert-major packed buffer + 首步 SIGSEGV 定位为 base-weight backward 错用路由表,现改用 expert-major packed buffer; + down 梯度为零进一步定位为 route-weighted `grad_output` 被 gate/up backward 临时区覆盖, + 现以独立只读快照保留,并用 NUMA 子线程池并行计算 AMX BF16 基座梯度 ## 1. 开发前先确认 @@ -150,6 +152,12 @@ router weight,因而同时保证 down projection 基座梯度的权重语义 6. `lora_rank=0` 时 scaling 明确为 0,避免纯 Full FT 产生 Inf。 7. 基座梯度读取 expert-major packed input/grad-output,不把 `m_local_pos_cache[token][route]` 误当作 expert token 列表。 +8. down 基座梯度读取独立的 route-weighted `grad_output` 快照;工作用 + `grad_output_bf16_ptr_` 后续可继续被 gate/up grad-input 路径复用。 +9. AMX BF16 基座梯度按 expert、projection 和 `32x32` 输出 tile 提交给当前 NUMA + subpool。gate/up 共用输入 tile,AMX BF16 乘法以 FP32 累加,每个任务独占输出区域。 +10. `lora_rank=0` 在完成 gate/up 基座 grad-input 后直接跳过 LoRA remainder,避免纯 Full FT + 进入零秩 LoRA 临时区路径。 ### 6.2 `operators/moe-sft-tp.hpp` @@ -165,8 +173,10 @@ router weight,因而同时保证 down projection 基座梯度的权重语义 ## 7. 验证顺序 -当前已通过 `clang-format --dry-run --Werror` 和 AMX/CUDA Release -`build_ext --inplace`。但当前 Kllama 环境中的 `kt_kernel_ext` 已 stripped, +当前已通过 `clang-format --dry-run --Werror`、AMX/CUDA Release +`build_ext --inplace`,以及 AMX BF16 `32x32x37` 分块乘法与标量 BF16 参考的独立校验 +(`max_abs_error=0`)。端到端短训练仍须确认 gate/up/down 三组梯度均非零且有限。 +但当前 Kllama 环境中的 `kt_kernel_ext` 已 stripped, 普通 GDB 只能得到地址或有限符号;定位 SIGSEGV 前应使用相同 Python 环境重新构建 并安装带符号的 `RelWithDebInfo` 版本: diff --git a/kt-kernel/operators/amx/sft_moe.hpp b/kt-kernel/operators/amx/sft_moe.hpp index efb7d6f54..aa637cc4c 100644 --- a/kt-kernel/operators/amx/sft_moe.hpp +++ b/kt-kernel/operators/amx/sft_moe.hpp @@ -576,10 +576,15 @@ class AMX_SFT_MOE_TP : public BaseMOE { // BF16 buffer for scattered grad_output (before quantization to BufferA) std::vector grad_output_bf16_ptr_; // [expert_num] + // Immutable copy of the route-weighted grad_output used by full-weight down_proj gradients. + // grad_output_bf16_ptr_ is reused as gate/up grad_input scratch later in backward. + std::vector base_grad_output_bf16_ptr_; // [expert_num] + // Backward buffer pools void* backward_ba_pool_ = nullptr; void* backward_bc_pool_ = nullptr; void* grad_output_bf16_pool_ = nullptr; + void* base_grad_output_bf16_pool_ = nullptr; void* backward_pool_ = nullptr; size_t backward_pool_bytes_ = 0; @@ -587,6 +592,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { size_t backward_ba_pool_bytes_ = 0; size_t backward_bc_pool_bytes_ = 0; size_t grad_output_bf16_pool_bytes_ = 0; + size_t base_grad_output_bf16_pool_bytes_ = 0; // LoRA gradient computation pools (FP32, used in bwd_down_lora_precompute and grad computation) float* lora_grad_out_pool_ = nullptr; // [max_len * num_experts_per_tok * hidden_size] @@ -821,6 +827,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { required += round_up(backward_ba_pool_bytes_, kAmxAlignment); required += round_up(backward_bc_pool_bytes_, kAmxAlignment); required += round_up(grad_output_bf16_pool_bytes_, kAmxAlignment); + required += round_up(base_grad_output_bf16_pool_bytes_, kAmxAlignment); required += round_up(lora_grad_out_pool_bytes_, kAmxAlignment); required += round_up(lora_inter_proj_pool_bytes_, kAmxAlignment); required += round_up(lora_grad_times_b_pool_bytes_, kAmxAlignment); @@ -853,6 +860,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { assign(&backward_ba_pool_, backward_ba_pool_bytes_); assign(&backward_bc_pool_, backward_bc_pool_bytes_); assign(&grad_output_bf16_pool_, grad_output_bf16_pool_bytes_); + assign(&base_grad_output_bf16_pool_, base_grad_output_bf16_pool_bytes_); assign((void**)&lora_grad_out_pool_, lora_grad_out_pool_bytes_); assign((void**)&lora_inter_proj_pool_, lora_inter_proj_pool_bytes_); @@ -1919,7 +1927,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { const int H = config_.hidden_size; const int I = config_.intermediate_size; const int F = full_intermediate_size; - int activated_expert = cache.activated_expert_cache; + const int activated_expert = cache.activated_expert_cache; if (F < I) { throw std::runtime_error("full_intermediate_size must be at least the TP-local intermediate_size"); @@ -1929,15 +1937,163 @@ class AMX_SFT_MOE_TP : public BaseMOE { auto* gup_ptr = static_cast(grad_up_proj); // TP slice of [E, F, H] auto* gdp = static_cast(grad_down_proj); // TP slice of [E, H, F] + std::vector expert_offsets(activated_expert); + size_t token_offset = 0; + for (int task_id = 0; task_id < activated_expert; task_id++) { + expert_offsets[task_id] = token_offset; + int expert_idx = cache.m_expert_id_map_cache[task_id]; + token_offset += cache.m_local_num_cache[expert_idx]; + } + + if constexpr (std::is_same_v && amx::AMX_AVAILABLE) { + constexpr int TILE_M = amx::GemmKernel224BF::M_STEP; + constexpr int TILE_N = amx::GemmKernel224BF::N_STEP; + constexpr int TILE_K = amx::GemmKernel224BF::K_STEP; + static_assert(TILE_M == 32 && TILE_N == 32 && TILE_K == 32, + "base-weight gradient tile packing assumes 32x32x32 AMX BF16 tiles"); + + const int i_tiles = (I + TILE_M - 1) / TILE_M; + const int h_tiles = (H + TILE_N - 1) / TILE_N; + const int tiles_per_projection = i_tiles * h_tiles; + const int tasks_per_expert = tiles_per_projection * 2; // fused gate/up plus down + const int total_tasks = activated_expert * tasks_per_expert; + auto pool = config_.pool->get_subpool(tp_part_idx); + + pool->do_work_stealing_job( + total_tasks, [](int _) { T::config(); }, + [&, i_tiles, h_tiles, tiles_per_projection, tasks_per_expert](int task_id) { + const int expert_task = task_id / tasks_per_expert; + const int local_task = task_id % tasks_per_expert; + const bool do_down = local_task >= tiles_per_projection; + const int tile_id = do_down ? local_task - tiles_per_projection : local_task; + const int expert_idx = cache.m_expert_id_map_cache[expert_task]; + const int m = cache.m_local_num_cache[expert_idx]; + if (m == 0) return; + + const size_t pos_start = expert_offsets[expert_task]; + alignas(64) ggml_bf16_t a_tile[TILE_M * TILE_K]; + alignas(64) ggml_bf16_t b_tile[TILE_N * TILE_K]; + alignas(64) float c0[TILE_M * TILE_N]; + alignas(64) float c1[TILE_M * TILE_N]; + + if (!do_down) { + const int i_tile = tile_id / h_tiles; + const int h_tile = tile_id % h_tiles; + const int i_start = i_tile * TILE_M; + const int h_start = h_tile * TILE_N; + const int i_count = std::min(TILE_M, I - i_start); + const int h_count = std::min(TILE_N, H - h_start); + const ggml_bf16_t* input = m_local_input_ptr_[expert_idx]; + + for (int k_start = 0; k_start < m; k_start += TILE_K) { + const int k_count = std::min(TILE_K, m - k_start); + std::memset(a_tile, 0, sizeof(a_tile)); + std::memset(b_tile, 0, sizeof(b_tile)); + + for (int row = 0; row < i_count; row++) { + for (int kk = 0; kk < k_count; kk++) { + a_tile[row * TILE_K + kk] = grad_gate_output_[(pos_start + k_start + kk) * I + i_start + row]; + } + } + for (int col = 0; col < h_count; col++) { + for (int kk = 0; kk < k_count; kk++) { + b_tile[col * TILE_K + kk] = input[(size_t)(k_start + kk) * H + h_start + col]; + } + } + amx::transpose_16x16_32bit(reinterpret_cast<__m512i*>(b_tile)); + amx::transpose_16x16_32bit(reinterpret_cast<__m512i*>(b_tile + amx::GemmKernel224BF::TILE_N * TILE_K)); + + T::load_b(b_tile, TILE_K * sizeof(ggml_bf16_t)); + T::load_a(a_tile, TILE_K * sizeof(ggml_bf16_t)); + if (k_start == 0) { + T::clean_c(); + } else { + T::load_c(c0, TILE_N * sizeof(float)); + } + T::run_tile(); + T::store_c(c0, TILE_N * sizeof(float)); + + std::memset(a_tile, 0, sizeof(a_tile)); + for (int row = 0; row < i_count; row++) { + for (int kk = 0; kk < k_count; kk++) { + a_tile[row * TILE_K + kk] = grad_up_output_[(pos_start + k_start + kk) * I + i_start + row]; + } + } + T::load_a(a_tile, TILE_K * sizeof(ggml_bf16_t)); + if (k_start == 0) { + T::clean_c(); + } else { + T::load_c(c1, TILE_N * sizeof(float)); + } + T::run_tile(); + T::store_c(c1, TILE_N * sizeof(float)); + } + + ggml_bf16_t* gate_dst = ggp + (size_t)expert_idx * F * H; + ggml_bf16_t* up_dst = gup_ptr + (size_t)expert_idx * F * H; + for (int row = 0; row < i_count; row++) { + for (int col = 0; col < h_count; col++) { + gate_dst[(size_t)(i_start + row) * H + h_start + col] = GGML_FP32_TO_BF16(c0[row * TILE_N + col]); + up_dst[(size_t)(i_start + row) * H + h_start + col] = GGML_FP32_TO_BF16(c1[row * TILE_N + col]); + } + } + return; + } + + const int h_tile = tile_id / i_tiles; + const int i_tile = tile_id % i_tiles; + const int h_start = h_tile * TILE_M; + const int i_start = i_tile * TILE_N; + const int h_count = std::min(TILE_M, H - h_start); + const int i_count = std::min(TILE_N, I - i_start); + const ggml_bf16_t* grad_output = base_grad_output_bf16_ptr_[expert_idx]; + const ggml_bf16_t* intermediate = cache.intermediate_cache + pos_start * I; + + for (int k_start = 0; k_start < m; k_start += TILE_K) { + const int k_count = std::min(TILE_K, m - k_start); + std::memset(a_tile, 0, sizeof(a_tile)); + std::memset(b_tile, 0, sizeof(b_tile)); + for (int row = 0; row < h_count; row++) { + for (int kk = 0; kk < k_count; kk++) { + a_tile[row * TILE_K + kk] = grad_output[(size_t)(k_start + kk) * H + h_start + row]; + } + } + for (int col = 0; col < i_count; col++) { + for (int kk = 0; kk < k_count; kk++) { + b_tile[col * TILE_K + kk] = intermediate[(size_t)(k_start + kk) * I + i_start + col]; + } + } + amx::transpose_16x16_32bit(reinterpret_cast<__m512i*>(b_tile)); + amx::transpose_16x16_32bit(reinterpret_cast<__m512i*>(b_tile + amx::GemmKernel224BF::TILE_N * TILE_K)); + + T::load_b(b_tile, TILE_K * sizeof(ggml_bf16_t)); + T::load_a(a_tile, TILE_K * sizeof(ggml_bf16_t)); + if (k_start == 0) { + T::clean_c(); + } else { + T::load_c(c0, TILE_N * sizeof(float)); + } + T::run_tile(); + T::store_c(c0, TILE_N * sizeof(float)); + } + + ggml_bf16_t* down_dst = gdp + (size_t)expert_idx * H * F; + for (int row = 0; row < h_count; row++) { + for (int col = 0; col < i_count; col++) { + down_dst[(size_t)(h_start + row) * F + i_start + col] = GGML_FP32_TO_BF16(c0[row * TILE_N + col]); + } + } + }, + nullptr); + return; + } + for (int task_id = 0; task_id < activated_expert; task_id++) { int expert_idx = cache.m_expert_id_map_cache[task_id]; int m = cache.m_local_num_cache[expert_idx]; if (m == 0) continue; - int pos_start = 0; - for (int prev_id = 0; prev_id < task_id; prev_id++) { - pos_start += cache.m_local_num_cache[cache.m_expert_id_map_cache[prev_id]]; - } + const size_t pos_start = expert_offsets[task_id]; // Allocate FP32 accumulators from forward pool (safe during backward) float* acc_gate = static_cast(forward_pool_); // [I, H] @@ -1953,7 +2109,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { const ggml_bf16_t* gate_grad_row = grad_gate_output_ + (size_t)(pos_start + t) * I; const ggml_bf16_t* up_grad_row = grad_up_output_ + (size_t)(pos_start + t) * I; const ggml_bf16_t* inter_row = cache.intermediate_cache + (size_t)(pos_start + t) * I; - const ggml_bf16_t* grad_out_row = grad_output_bf16_ptr_[expert_idx] + (size_t)t * H; + const ggml_bf16_t* grad_out_row = base_grad_output_bf16_ptr_[expert_idx] + (size_t)t * H; // gate_proj grad: [I, H] += grad_gate_out[t]^T @ input[t] for (int i = 0; i < I; i++) { @@ -2835,6 +2991,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { // BF16 buffer for scattered grad_output grad_output_bf16_pool_bytes_ = safe_alloc_tokens * config_.hidden_size * sizeof(ggml_bf16_t) + align_overhead; + base_grad_output_bf16_pool_bytes_ = grad_output_bf16_pool_bytes_; // LoRA gradient computation FP32 pools (used in bwd_down_lora_precompute and grad computation) // Total tokens across all activated experts = safe_alloc_tokens @@ -2867,6 +3024,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { backward_ba_pool_bytes_ = 0; backward_bc_pool_bytes_ = 0; grad_output_bf16_pool_bytes_ = 0; + base_grad_output_bf16_pool_bytes_ = 0; backward_bb_pool_bytes_ = 0; lora_grad_out_pool_bytes_ = 0; lora_inter_proj_pool_bytes_ = 0; @@ -3007,6 +3165,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { grad_intermediate_bc_.resize(config_.expert_num); grad_gate_up_bc_.resize(config_.expert_num); grad_output_bf16_ptr_.resize(config_.expert_num); + base_grad_output_bf16_ptr_.resize(config_.expert_num); // Resize vectors - backward BufferB (transposed base weights) gate_backward_bb_.resize(config_.expert_num); @@ -3097,6 +3256,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { // BF16 pointer: will be assigned in backward grad_output_bf16_ptr_[i] = nullptr; + base_grad_output_bf16_ptr_[i] = nullptr; } // ===================================================== @@ -4297,6 +4457,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { char* backward_ba_ptr = (char*)backward_ba_pool_; char* backward_bc_ptr = (char*)backward_bc_pool_; char* grad_output_bf16_ptr = (char*)grad_output_bf16_pool_; + char* base_grad_output_bf16_ptr = (char*)base_grad_output_bf16_pool_; for (int task_id = 0; task_id < activated_expert; task_id++) { int expert_idx = m_expert_id_map_[task_id]; @@ -4318,6 +4479,11 @@ class AMX_SFT_MOE_TP : public BaseMOE { // Allocate BF16 buffer for scattered grad_output grad_output_bf16_ptr_[expert_idx] = (ggml_bf16_t*)grad_output_bf16_ptr; grad_output_bf16_ptr += align64(local_max_m * config_.hidden_size * sizeof(ggml_bf16_t)); + + // Preserve the route-weighted upstream gradient for down_proj full-weight gradients. + // The working grad_output buffer is overwritten by backward_gate_up_amx(). + base_grad_output_bf16_ptr_[expert_idx] = (ggml_bf16_t*)base_grad_output_bf16_ptr; + base_grad_output_bf16_ptr += align64(local_max_m * config_.hidden_size * sizeof(ggml_bf16_t)); } // NOTE: no full-buffer memset here; grad_intermediate_ is overwritten by to_mat() for active tokens. @@ -4377,6 +4543,17 @@ class AMX_SFT_MOE_TP : public BaseMOE { }); } + // Keep an immutable copy before grad_output_bf16_ptr_ is reused as gate/up grad_input scratch. + if (sft_config_.full_weight_grad) { + direct_or_pool(activated_expert, [this](int task_id) { + int expert_idx = m_expert_id_map_[task_id]; + int num_tokens = m_local_num_[expert_idx]; + if (num_tokens == 0) return; + std::memcpy(base_grad_output_bf16_ptr_[expert_idx], grad_output_bf16_ptr_[expert_idx], + static_cast(num_tokens) * config_.hidden_size * sizeof(ggml_bf16_t)); + }); + } + // ===================================================== // Step 3: Quantize scattered grad_output to BufferA // ===================================================== @@ -5031,7 +5208,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { ggml_bf16_t* grad_up_b = (ggml_bf16_t*)grad_up_lora_b; assert(backward_weights_prepared_); - if (gate_lora_a_ != nullptr && gate_lora_b_ != nullptr) { + if (lora_rank_ > 0 && gate_lora_a_ != nullptr && gate_lora_b_ != nullptr) { prepare_lora_backward_weights(); } @@ -5216,7 +5393,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { // } // Skip all LoRA computation when SkipLoRA is true - if (SkipLoRA || gate_lora_a_ == nullptr || gate_lora_b_ == nullptr) { + if (SkipLoRA || lora_rank_ <= 0 || gate_lora_a_ == nullptr || gate_lora_b_ == nullptr) { return; } From 6aa18620d2b286d5417ace282d83e697c021287f Mon Sep 17 00:00:00 2001 From: illu Date: Mon, 13 Jul 2026 07:54:16 +0000 Subject: [PATCH 05/20] [docs](kt-kernel): document Full FT fork changes and debug history --- kt-kernel/AGENTS.md | 648 ++++++++++++++++++++++++++++++++------------ 1 file changed, 471 insertions(+), 177 deletions(-) diff --git a/kt-kernel/AGENTS.md b/kt-kernel/AGENTS.md index 2648731b7..c404b32f4 100644 --- a/kt-kernel/AGENTS.md +++ b/kt-kernel/AGENTS.md @@ -1,78 +1,197 @@ -# AGENTS.md: kt-kernel Full FT 开发指南 +# AGENTS.md:kt-kernel MoE Full FT 开发、差异与调试记录 -本文件作用于 `kt-kernel/` 及其全部子目录,用于指导 Agent 开发和验证 MoE -Full Fine-Tuning(Full FT)功能。重点是专家基座权重的梯度、optimizer 更新 -和 AMX 权重重新量化,不讨论普通 GPU 参数训练。 +本文件作用于 `kt-kernel/` 及其全部子目录,用于说明个人 Full Fine-Tuning +(Full FT)分支相对官方仓库的真实代码差异,并保留从“专家权重不更新”到 +“SIGSEGV、down 梯度为零、性能过低”这一系列问题的试错记录。 -当前开发快照: +这里的 Full FT 专指 KTransformers CPU/AMX MoE expert 基座权重 +`gate_proj`、`up_proj`、`down_proj` 的训练,不把普通 GPU attention、router、 +embedding 等参数的更新当作 expert Full FT 已生效的证据。 -- 仓库:`Illumination111/ktransformers-fullFT_development` -- Full FT 基线提交:`e99b5e1`,基于上游 `8e46e58` -- Full FT 来源:poryfly 的 `239bac5`、`25fa2fd` -- 当前状态:TP 子核配置、梯度分片/stride、临时区容量和逐步清零已修复;GDB 已将 - 首步 SIGSEGV 定位为 base-weight backward 错用路由表,现改用 expert-major packed buffer; - down 梯度为零进一步定位为 route-weighted `grad_output` 被 gate/up backward 临时区覆盖, - 现以独立只读快照保留,并用 NUMA 子线程池并行计算 AMX BF16 基座梯度 +## 1. 仓库、基线与比较口径 -## 1. 开发前先确认 +### 1.1 分支关系 -1. 在 `ktransformers` 中运行 `git status --short --branch`,不要覆盖已有改动。 -2. 使用符号名定位代码,不依赖本文行号。 -3. 先确认目标是 `full`、`hybrid` 还是 `lora`: - - `full`:训练 expert 基座权重,通常 `lora_rank=0`。 - - `hybrid`:同时训练基座权重和 LoRA。 - - `lora`:不得计算或更新基座权重。 -4. 修改前记录 TP 数量、NUMA 数量、`E/H/F` 和权重 dtype。 +截至 2026-07-13,本地比较对象为: -## 2. Full FT 数据契约 +- 官方仓库:`kvcache-ai/ktransformers`,remote 名为 `origin`。 +- 个人仓库:`Illumination111/ktransformers-fullFT_development`,remote 名为 + `development`。 +- 个人开发分支:`fullft-development`,当前提交 `5ce0767`。 +- 共同基线:官方提交 `8e46e58`。 +- Full FT 初始快照:`e99b5e1`,来源包含 poryfly 的 `239bac5`、`25fa2fd`。 +- 后续修复:`2d81e86`、`20f645c`、`5ce0767`。 -完整链路: +官方 `origin/main` 当时已前进到 `7c021b4`,比共同基线多 3 个与本轮 Full FT +基本无关的提交: ```text -KTConfig - -> wrapper 创建 CPU BF16 nn.Parameter - -> forward 将 Parameter 挂入 KTMoEFunction - -> C++ backward 写 grad_*_proj_buf - -> autograd 返回基座梯度 - -> optimizer 更新 *_proj_buf - -> update_base_weights 重新量化到 AMX BufferB +79b265b normalize compressed RAWINT4 weights +cb9f47d detect bound ports before launch +7c021b4 bump sglang submodule ``` -权威数据必须唯一: +因此,统计“个人做了什么”时必须使用共同基线或三点 diff: -- `gate/up/down_proj_buf`:optimizer 可见的 CPU BF16 基座权重。 -- `grad_gate/up/down_proj_buf`:C++ backward 的输出。 +```bash +git merge-base origin/main fullft-development +git diff --shortstat origin/main...fullft-development +git diff --numstat origin/main...fullft-development +``` + +不要直接把 `git diff origin/main fullft-development` 的全部结果都归到个人修改; +那样会把个人分支尚未合入的 3 个官方提交也统计为差异。 + +### 1.2 修改规模 + +个人 Full FT 分支相对共同基线 `8e46e58`: + +| 口径 | 文件数 | 新增 | 删除 | 说明 | +|---|---:|---:|---:|---| +| 当前工作树(含本次扩写后的文档) | 16 | 1581 | 160 | 全部位于 `kt-kernel/` | +| 仅生产代码 | 15 | 1036 | 160 | 总代码 churn 约 1196 行 | +| Python SFT 接入 | 11 | 658 | 143 | 约占生产代码改动的三分之二 | +| C++ config/binding/AMX/TP | 4 | 378 | 17 | 约占生产代码改动的三分之一 | + +应区分两个阶段: + +1. `e99b5e1` 是完整 Full FT 功能接入,涉及 15 个代码文件,`+809/-155`, + 修改面较广。 +2. 后续 3 个 debug/fix 提交的生产代码只修改两个文件,合计约 `+255/-33`: + `operators/amx/sft_moe.hpp` 和 `operators/moe-sft-tp.hpp`。 + +提交 `5ce0767` 时旧版本文档为 251 行;本次仅扩写说明和试错记录,没有改变上述 +15 个生产代码文件的统计。 + +因此,“个人 Full FT 分支总体改动”属于中等偏大;“本轮 SIGSEGV/down 梯度/ +AMX 性能 debug”则高度集中在两个 `.hpp` 文件。 + +## 2. 官方代码与个人 Full FT 代码的行为差异 + +### 2.1 官方基线的限制 + +官方基线的 SFT expert 路径以 LoRA 为主要训练对象: + +- 没有 `full/hybrid/lora` 的明确 KT train-mode 映射。 +- `lora_rank` 必须为正数,纯 Full FT 的 `lora_rank=0` 不成立。 +- expert 原始权重只用于加载和 AMX forward,不是 optimizer 可见的 + `nn.Parameter`。 +- Python autograd 不把三个 expert 基座权重作为输入,也不返回对应梯度。 +- C++ `MOESFTConfig` 没有 `full_weight_grad` 和三个基座梯度指针。 +- C++ backward 不会把 gate/up/down 基座梯度写回 Python tensor。 +- optimizer 更新后没有把 BF16 expert 基座权重重新量化到 AMX BufferB 的闭环。 + +所以,仅在 LLaMA-Factory 中设置 `finetuning_type: full`,不能证明 KT CPU experts +正在 Full FT;attention/router 等 GPU 参数仍可能让 loss 下降。 + +### 2.2 个人分支建立的完整数据链路 + +```text +LLaMA-Factory finetuning_type + -> KTConfig.kt_train_mode / kt_full_weight_grad + -> wrapper 创建 CPU BF16 gate/up/down nn.Parameter + -> KTMoEFunction.forward 显式接收三个 Parameter + -> C++ backward 写 grad_gate/up/down_proj_buf + -> KTMoEFunction.backward 返回三个基座梯度 + -> Trainer optimizer 注入并更新三个 *_proj_buf + -> optimizer.step 后标记 _base_weights_dirty + -> 下一次 forward 调用 update_base_weights + -> BF16 基座权重重新量化到 AMX BufferB +``` + +权威数据定义: + +- `gate_proj_buf`、`up_proj_buf`、`down_proj_buf`:optimizer 可见的 CPU BF16 + expert 基座权重。 +- `grad_gate_proj_buf`、`grad_up_proj_buf`、`grad_down_proj_buf`:C++ backward + 写入、autograd 返回的基座梯度。 - AMX 量化权重:forward 使用的派生副本,optimizer 后必须重建。 -- HF expert 原权重只是占位,不得作为 Full FT 的更新依据。 +- HF model tree 中的 expert 权重:清理为 zero-storage placeholder,不再是 Full FT + 的权威副本,也不能用于训练前后权重比较。 + +### 2.3 模式语义 + +- `full`:训练 expert 基座权重,通常 `lora_rank=0`。 +- `hybrid`:同时训练 expert 基座权重和 LoRA,允许 `lora_rank>0`。 +- `lora`:只训练 LoRA,不得计算、同步或更新 expert 基座权重。 + +修改任何共享路径前必须分别确认三种模式,不能用 Full FT 修复破坏 LoRA。 -## 3. 文件职责 +## 3. 逐文件说明:个人分支具体修改了什么 -路径均相对本文件所在的 `kt-kernel/`。 +路径均相对 `kt-kernel/`。 -| 路径 | 职责 | +### 3.1 Python 训练接入层 + +| 文件 | 相对官方的主要修改 | +|---|---| +| `python/sft/config.py` | 增加 `kt_train_mode`、`kt_full_weight_grad`;读取 `ACCELERATE_KT_TRAIN_MODE`,将 `full/hybrid` 映射为基座梯度开启。 | +| `python/experts.py` | `KTMoEWrapper` 和 SFT wrapper factory 接受并传递 `full_weight_grad`。 | +| `python/sft/wrapper.py` | 将 LLaMA-Factory `finetuning_type` 映射到 KT train mode;把 `lora_rank=0` 当作合法纯 Full FT;初始化三个权威 BF16 Parameter/grad buffer;释放重复权重;把 HF expert 权重替换为 zero-storage placeholder。 | +| `python/sft/base.py` | 保存 `_full_weight_grad` 和 `_base_weights_dirty`;创建 `[E,F,H]` gate/up 与 `[E,H,F]` down Parameter/grad buffer;允许 Full FT 在无 LoRA 时运行;定义 `update_base_weights()` 接口;避免 `lora_rank=0` 除零。 | +| `python/sft/layer.py` | Full FT 时强制进入 autograd 路径;将三个基座 Parameter 传入 `KTMoEFunction`;optimizer 后发现 dirty 权重时触发 requant;Full FT 时保留 router 梯度;兼容 transformers v5 TopKRouter/GLM4 路由。 | +| `python/sft/autograd.py` | forward 增加三个 expert 基座 Parameter 输入;backward 返回 C++ 写入的 gate/up/down 梯度,使 PyTorch 给 Parameter 建立 `.grad`。 | +| `python/sft/lora.py` | 拆分 LoRA 参数与 Full-FT 基座参数收集;把 48 层 × 3 投影共 144 个 Parameter 注入 optimizer;纯 Full FT 跳过 LoRA buffer;分布式时同步基座梯度;optimizer 后标记基座权重 dirty。 | +| `python/sft/amx.py` | 将 Full-FT 开关和三个梯度 data pointer 写入 C++ config/backward task;`lora_rank=0` 时传空 LoRA 指针;增加 `update_base_weights()`,优先复用 C++ 对象并重新量化,缺少 binding 时才完整重建。 | +| `python/sft/weights.py` | 明确 `*_proj_buf` 为 Full FT 权威权重;清理 model tree 中的冗余 expert 参数并打 `_kt_zero_storage` 标记,避免重复计数和重复占内存。 | +| `python/sft/arch.py` | 增加 GLM4 MoE 架构识别;这属于同期兼容性修改,不是本次 Qwen3 Full FT bugfix 的核心。 | +| `python/sft/__init__.py` | 导出新增的 Full-FT 参数收集和相关 SFT API。 | + +### 3.2 C++ 配置与绑定 + +| 文件 | 相对官方的主要修改 | |---|---| -| `python/sft/config.py` | 将 `full/hybrid` 映射为 `kt_full_weight_grad=True` | -| `python/sft/base.py` | 创建 `*_proj_buf`、`grad_*_proj_buf` | -| `python/sft/wrapper.py` | 包装 MoE、确定权威权重 | -| `python/sft/layer.py` | 将三个基座 Parameter 传入 autograd;dirty 时 requant | -| `python/sft/autograd.py` | 返回 C++ 写入的基座梯度 | -| `python/sft/lora.py` | optimizer 参数注入、分布式同步、dirty 标记 | -| `python/sft/amx.py` | 设置 C++ config、传梯度指针、触发重新量化 | -| `operators/common.hpp` | `MOESFTConfig.full_weight_grad` | -| `operators/moe-tp.hpp` | 通用 TP 切分;当前会把 SFT config 切成基类 | -| `operators/moe-sft-tp.hpp` | SFT TP 调度和全局梯度 buffer 分片 | -| `operators/amx/sft_moe.hpp` | NUMA 子核 backward 和基座梯度计算 | -| `ext_bindings.cpp` | Python/C++ 参数和 task 绑定 | - -## 4. TP 张量布局 +| `operators/common.hpp` | 在 `MOESFTConfig` 中增加 `full_weight_grad` 和 `grad_gate/up/down_proj` 三个零拷贝指针。 | +| `ext_bindings.cpp` | backward binding 接收并转发三个基座梯度指针;向 Python 暴露 Full-FT config 字段;增加 `set_base_weight_pointers()`,支持 optimizer 后复用既有 C++ MoE 对象并 requant。 | + +### 3.3 两个核心 `.hpp` + +#### `operators/moe-sft-tp.hpp` + +该文件负责顶层 TP 调度和完整梯度 tensor 的切片: + +1. 子核创建后显式调用 `set_full_weight_grad()`,补回派生 config 被 + `GeneralMOEConfig` slicing 丢失的字段。 +2. 传播逻辑移到 `if constexpr (!kSkipLoRA)` 外,确保纯 Full FT 的 + `lora_rank=0` 仍启用基座梯度。 +3. backward dispatch 前按 TP offset 计算 gate/up/down 三个 slice 指针。 +4. 将完整 `F` 传给子核作为 global stride,而不是误用当前 TP 的本地 `I`。 +5. 检查所有 TP slice 对完整 `F` 的覆盖,无重叠、无缺口。 +6. dispatch 前统一并行清零三组完整梯度,避免跨 step 残留和多个 TP 子核竞争清零。 +7. 空指针必须先判断再做 offset,避免对 `nullptr` 做未定义指针运算。 + +#### `operators/amx/sft_moe.hpp` + +该文件负责每个 NUMA/TP 子核的真实 backward 和基座梯度计算: + +1. 保存并更新 `sft_config_.full_weight_grad`。 +2. `backward_base_weight_grad` 同时接收本地 `I` 和完整 `F`;本地计算用 `I`, + 写回 expert/row stride 用 `F`。 +3. gate/up/down 三组指针全部有效时才计算基座梯度。 +4. 扩大/复用工作区,保证基座 FP32 累加需要的容量。 +5. `lora_rank=0` 时 scaling 为 0,完成 gate/up grad-input 后跳过 LoRA remainder, + 不进入零秩临时区路径。 +6. 不再把 token-major 路由表当作 expert-major token list;直接读取 backward 已 + 打包好的 `m_local_input_ptr_[expert]` 和 expert-major grad-output。 +7. 为 down 梯度保存独立、只读、已经乘过 router weight 的 `grad_output` 快照; + gate/up grad-input 可以继续复用工作 buffer,但不能覆盖 down 所需的 dY。 +8. 使用 AMX BF16 `32x32` output tile、FP32 累加计算基座梯度。 +9. 任务按 expert、projection、tile 投递给当前 NUMA subpool,而不是在每个 NUMA + 节点只用单核执行三重标量循环。 +10. gate/up 共用输入 tile,每个任务写独占输出区域,避免锁和写竞争。 + +暂不修改通用 `operators/moe-tp.hpp`。将其改成保留派生 config 的泛型重构会影响 +所有 MoE backend,应另开改动并做完整回归。 + +## 4. TP 梯度布局契约 定义: - `E`:expert 数量。 - `H`:hidden size。 -- `F`:完整 `intermediate_size`。 -- `I`:当前 TP 子核的本地 `intermediate_size`。 -- `tp_offset`:当前 TP 在完整 `F` 维上的起点。 +- `F`:完整 intermediate size。 +- `I`:当前 TP 子核的本地 intermediate size。 +- `tp_offset`:当前 TP slice 在完整 `F` 维的起点。 完整梯度布局: @@ -81,104 +200,264 @@ gate/up: [E, F, H] down: [E, H, F] ``` -TP wrapper 传给子核的起始指针: +顶层 TP wrapper 传给子核的起始指针: ```text gate/up = base + tp_offset * H down = base + tp_offset ``` -子核循环和 FP32 累加使用本地 `I`,但写回完整张量时必须使用 `F` -作为 expert/row stride: +子核写回公式: ```text gate/up: expert * F * H + local_i * H + h down: expert * H * F + h * F + local_i ``` -所有 TP 必须写入互不重叠的区域;对空指针不得做偏移运算。 +易错点: + +- 计算循环范围是本地 `I`,但 expert/row stride 必须是完整 `F`。 +- 所有 TP 必须写入互不重叠的区域。 +- `expert_id=0` 可能掩盖错误 expert stride,参考测试必须包含 `expert_id>0`。 +- 梯度清零只能由顶层统一完成,不能让多个 TP 子核同时清相同完整 buffer。 + +## 5. 修改与试错时间线 + +本节保留失败路径,因为这些失败揭示了仅看 loss、requant 日志或单步权重抽检会得出 +错误结论。 + +### 5.1 阶段 0:官方路径不能形成 expert Full FT 闭环 + +最初仅设置 LLaMA-Factory `finetuning_type=full`,但 KT expert 权重不是 optimizer +可见 Parameter,C++ 也没有基座梯度输出。loss 变化最多说明其他参数在训练,不能说明 +CPU experts 更新。 -## 5. 历史根因与失败基线 +结论:必须建立 Parameter → C++ grad → autograd → optimizer → requant 的完整闭环。 -`full_weight_grad=True` 已正确到达顶层 `TP_MOE_SFT`,但 -`TP_MOE` 创建子核时执行: +### 5.2 阶段 1:`e99b5e1` 接入 Full FT,但专家权重仍不更新 -```cpp -GeneralMOEConfig tp_config = config; +`e99b5e1` 加入 Python/C++ Full-FT 数据链路。随后运行: + +```text +20260710_103230_1gpu_AMX_BF16_FULLFT_EXPERTONLY +``` + +测试已冻结非 expert 参数,并向 optimizer 注入 144 个 expert 基座 Parameter,但结果为: + +- 15 步 `grad_norm=0`。 +- step 0 到 step 15 抽检 `changed=0/12`。 +- `max_abs_delta=0`。 +- 有 720 次 requant 日志。 + +试错结论: + +- optimizer 参数数量正确,不等于 C++ 真正产生了梯度。 +- requant 被调用,不等于其输入权重发生了变化。 +- 问题继续下沉到 C++ 子核配置和 TP 梯度写回。 + +根因之一是顶层 `TP_MOE_SFT` 收到 `full_weight_grad=true` 后,创建子核时把派生配置 +赋给 `GeneralMOEConfig`,发生 slicing;子核重建 `MOESFTConfig` 后开关恢复为 false, +`backward_base_weight_grad` 被静默跳过。 + +### 5.3 阶段 2:`2d81e86` 修复开关/TP stride 后出现 SIGSEGV + +`2d81e86` 尝试修复: + +- 将 `full_weight_grad` 显式传播到子核。 +- 给每个 TP 传独立梯度 slice。 +- 使用完整 `F` 作为全局 stride。 +- 扩大临时区并统一清零。 + +但测试: + +```text +20260710_192456_1gpu_AMX_BF16_FULLFT_EXPERTONLY +``` + +在第一个训练 step 前后发生 SIGSEGV: + +- child `exitcode=-11`,原始日志明确记录 `Signal 11 (SIGSEGV)`。 +- 外层 `exit_code.txt=1` 只是 accelerate launcher 退出码。 +- 当时的自动 `summary.md` 错误写成“未检测到崩溃”,不能作为反证。 +- 该 run 没有 C++ backtrace,最初只能假设是 gradient index/stride 越界。 + +失败原因:修复了“是否计算”和“写到哪里”,但新实现为了计算基座梯度又错误理解了 +路由缓存布局。 + +### 5.4 阶段 3:加入带符号 GDB,定位错误路由表解释 + +为避免继续猜测,在 FFTtest runner 中增加 batch GDB,并用同一 Kllama Python 环境 +构建 `RelWithDebInfo` 扩展。日志: + +```text +20260713_105328_1gpu_AMX_BF16_FULLFT_EXPERTONLY +``` + +GDB 证据: + +- `m_local_pos_cache` 的真实布局是 `[token_idx][route_slot]`。 +- 每个 token 的 route 长度是 `k=8`。 +- 旧代码却按 `m_local_pos_cache[expert_idx][t]` 把它当作 expert token list。 +- 第一个 expert 执行到 `t=8` 即越界,读出垃圾 `tok_pos=81349952`。 +- 最终在读取 `input_row[0]` 时于当时的 `sft_moe.hpp:1968` 触发 SIGSEGV。 +- 两个 NUMA 子核都进入了相同错误路径。 + +这一步推翻了“只是 TP output stride 错误”的单一假设。崩溃发生在读取输入行,而不是 +最终写回梯度的位置。 + +### 5.5 阶段 4:`20f645c` 最小修复消除 SIGSEGV,但单步测试无效 + +`20f645c` 不再反查 token-major 路由表,改用 backward 已生成的 expert-major packed +buffer: + +```text +m_local_input_ptr_[expert_idx] +grad_output_bf16_ptr_[expert_idx] ``` -这会丢失 `MOESFTConfig.full_weight_grad`。子核随后通过 -`MOESFTConfig(GeneralMOEConfig)` 重建配置,字段回到默认 `false`, -因此 `backward_base_weight_grad` 被门控跳过。 - -只恢复该布尔值仍不够:当前 TP backward 把相同的完整梯度首地址传给 -所有子核,子核又按本地 `I` 写回,会导致覆盖和错误布局。 - -修复前的 `20260710_103230_1gpu_AMX_BF16_FULLFT_EXPERTONLY` run 已冻结 -其他权重并只向 optimizer 注入 144 个 expert 基座参数,但 15 步 -`grad_norm=0`,step 0 到 step 15 的 probe 为 `changed=0/12`、 -`max_abs_delta=0`。720 次 requant 不能证明权重更新。 - -修复后的 `20260710_192456_1gpu_AMX_BF16_FULLFT_EXPERTONLY` run 未进入 -第一个训练 step。`phase4/train.log` 记录训练子进程 PID 3296907 在 -`2026-07-10 19:27:33 +08:00` 以 `exitcode=-11` 退出,并明确报告 -`Signal 11 (SIGSEGV)`;外层 `accelerate` 的 `exit_code.txt=1` 只是 launcher -退出码。`phase4/log_analysis.txt` 正确识别到崩溃,但旧 `summary.md` 的 -“未检测到崩溃”是汇总脚本误判,不得作为反证。该 run 没有 C++ backtrace、 -源码行号或 core,因而“C++ 梯度索引越界”目前只是诊断假设,并非已确认根因。 - -带符号 GDB 的 `20260713_105328_1gpu_AMX_BF16_FULLFT_EXPERTONLY` run -确认根因位于 `backward_base_weight_grad`:`m_local_pos_cache` 的真实布局是 -`[token_idx][route_slot]`,内层长度为 `k=8`,旧代码却按 -`m_local_pos_cache[expert_idx][t]` 访问。首个 expert 执行到 `t=8` 时越界, -读出垃圾 `tok_pos=81349952`,最终在读取 `input_row[0]` 时于 -`operators/amx/sft_moe.hpp:1968` 触发 SIGSEGV。两个 NUMA 子核均进入同一错误路径。 -最小修复不再反查路由表,而是直接使用 backward 已恢复/生成的 -`m_local_input_ptr_[expert_idx]` 和 `grad_output_bf16_ptr_[expert_idx]`;后者已包含 -router weight,因而同时保证 down projection 基座梯度的权重语义正确。 - -## 6. 当前两文件实现 - -生产代码限制在两个文件,未修改 Python 接口。 - -### 6.1 `operators/amx/sft_moe.hpp` - -1. `set_full_weight_grad(bool)` 更新子核的 - `sft_config_.full_weight_grad`。 -2. `full_intermediate_size` 已传给 `backward_base_weight_grad`。 -3. 计算循环使用本地 `I`,写回 stride 使用完整 `F`。 -4. 门控要求开关为 true 且三个梯度指针均非空。 -5. `forward_pool_` 保证至少 `3 * I * H * sizeof(float)` 字节。 -6. `lora_rank=0` 时 scaling 明确为 0,避免纯 Full FT 产生 Inf。 -7. 基座梯度读取 expert-major packed input/grad-output,不把 - `m_local_pos_cache[token][route]` 误当作 expert token 列表。 -8. down 基座梯度读取独立的 route-weighted `grad_output` 快照;工作用 - `grad_output_bf16_ptr_` 后续可继续被 gate/up grad-input 路径复用。 -9. AMX BF16 基座梯度按 expert、projection 和 `32x32` 输出 tile 提交给当前 NUMA - subpool。gate/up 共用输入 tile,AMX BF16 乘法以 FP32 累加,每个任务独占输出区域。 -10. `lora_rank=0` 在完成 gate/up 基座 grad-input 后直接跳过 LoRA remainder,避免纯 Full FT - 进入零秩 LoRA 临时区路径。 - -### 6.2 `operators/moe-sft-tp.hpp` - -1. 子核创建后逐个调用 `set_full_weight_grad`。 -2. 传播逻辑位于 `if constexpr (!kSkipLoRA)` 外;纯 Full FT 的 - `lora_rank=0` 也必须执行。 -3. backward dispatch 前按第 4 节公式生成每个 TP 的三个梯度指针。 -4. 完整 `F` 传给子核作为全局 stride,并检查所有 TP slice 完整覆盖 `F`。 -5. TP dispatch 前统一并行清零三组完整梯度,避免跨 step 残留和 TP 清零竞争。 - -暂不修改通用 `operators/moe-tp.hpp`。保留派生 config 的泛型重构影响 -所有 MoE backend,应作为后续独立改动。 - -## 7. 验证顺序 - -当前已通过 `clang-format --dry-run --Werror`、AMX/CUDA Release -`build_ext --inplace`,以及 AMX BF16 `32x32x37` 分块乘法与标量 BF16 参考的独立校验 -(`max_abs_error=0`)。端到端短训练仍须确认 gate/up/down 三组梯度均非零且有限。 -但当前 Kllama 环境中的 `kt_kernel_ext` 已 stripped, -普通 GDB 只能得到地址或有限符号;定位 SIGSEGV 前应使用相同 Python 环境重新构建 -并安装带符号的 `RelWithDebInfo` 版本: +随后单步测试: + +```text +20260713_111729_1gpu_AMX_BF16_FULLFT_EXPERTONLY +``` + +结果: + +- GDB 显示进程正常退出,无 SIGSEGV。 +- backward 约 890.9 秒。 +- 权重抽检仍是 `changed=0/12`。 + +这里不能得出“梯度仍无效”的结论,因为只运行 1 step,而第一个 optimizer step 的 +learning rate 是 0。这个 run 证明了崩溃消失,但不能验证权重更新。 + +试错教训:至少需要 2 个 optimizer step,最好 3~5 步;必须同时记录每步 LR。 + +### 5.6 阶段 5:三步测试证明 gate/up 更新,但暴露 down 梯度为零 + +继续运行: + +```text +20260713_120636_1gpu_AMX_BF16_FULLFT_EXPERTONLY +``` + +结果: + +- 3 个 step 正常退出,无 SIGSEGV。 +- 权重抽检 `changed=5/12`,证明 Full-FT optimizer 链路总体已经能更新 expert。 +- gate 梯度 `48/48` 非零。 +- up 梯度 `48/48` 非零。 +- down 梯度 `0/48`,所以总计只有 `96/144` 非零。 +- backward 平均约 754.5 秒/step。 +- 梯度最大值达到 `4e16~9e16`,虽有限但数值明显异常。 + +为了区分“C++ buffer 有值但 autograd/optimizer 看不到”和“C++ 本身没有写值”,测试脚本 +增加了对 `grad_*_proj_buf` 与独立 `Parameter.grad` 的全量扫描。二者的非零数量和最大值 +一致,说明 down=0 发生在 C++ 计算链路,而不是 Python autograd 丢梯度。 + +进一步检查发现,down 基座梯度需要的 route-weighted dY 原本位于 +`grad_output_bf16_ptr_`,但 gate/up grad-input backward 随后复用了同一 buffer,down 计算 +时读到的内容已被覆盖。 + +失败方案:简单地把 `grad_output_bf16_ptr_` 当作长期只读 down dY。该指针只是工作 buffer, +生命周期不满足要求。 + +### 5.7 阶段 6:`5ce0767` 保存 down 快照并用 AMX/NUMA 并行 + +`5ce0767` 同时处理正确性和性能: + +- 在工作 buffer 被 gate/up 路径复用前,保存 route-weighted dY 的独立只读快照。 +- down 基座梯度只读取该快照。 +- gate/up/down 基座梯度改用 AMX BF16 `32x32` tile、FP32 累加。 +- 任务提交给每个 NUMA 节点已有的 subpool,而不是每节点单核标量循环。 +- 纯 Full FT `lora_rank=0` 跳过 LoRA remainder。 + +验证日志: + +```text +20260713_140101_1gpu_AMX_BF16_FULLFT_EXPERTONLY +``` + +结果: + +- 5/5 step 正常退出,无 SIGSEGV。 +- C++ grad buffer `144/144` 非零,`Parameter.grad` 也是 `144/144` 非零。 +- gate/up/down 分别为 `48/48`、`48/48`、`48/48`。 +- 权重抽检 `changed=9/12`;其中 down 抽检 `4/4` 全部变化。 +- 最大权重差约 `3.8147e-05`,没有非有限权重。 +- backward 从约 754.5 秒降到约 22.8 秒,约 33 倍加速、97% 降幅。 + +但该 run 只证明“链路连通、无非有限值、权重可更新”,不能证明数值完全正确: + +- 梯度最大值仍在 `1e16~1e17`。 +- loss 约 1~2 时,这个量级明显可疑。 +- AdamW 会归一化梯度,权重变化合理不能反推原始梯度合理。 + +### 5.8 探针本身造成的性能误判 + +全量梯度探针每个 step 扫描: + +- 144 个 C++ grad buffer。 +- 144 个独立 `Parameter.grad`。 +- 合计约 580 亿个 BF16 元素。 +- `isfinite`、`count_nonzero`、`abs().max()` 会重复遍历。 + +在 `20260713_140101` 中: + +- backward 已降到约 22.8 秒/step。 +- `step_other` 仍约 256.7 秒/step,占 61.8%。 +- 该 `other` 主要是 pre-optimizer 全量梯度扫描,不是 GDB。 + +因此诊断 run 与性能 run 必须分开: + +- 数值诊断:1~3 step,抽样或分阶段记录 C++ 中间量。 +- 性能测试:不使用 GDB、不扫描 expert 梯度/权重,只保留内存、显存和轻量 step timing。 +- 正式 TPS:固定 15 step,跳过前 5 个 warmup,用后 10 个 step 总 token/总时间计算。 + +## 6. 当前状态与仍未解决的问题 + +### 6.1 已确认解决 + +- `full_weight_grad` 能从 Python 到达顶层和所有 TP/NUMA 子核。 +- 三个 expert 基座 Parameter 能进入 optimizer。 +- TP gradient slice、global stride 和逐步清零已修复。 +- 错误解释 `m_local_pos_cache` 导致的 SIGSEGV 已修复。 +- down dY 被工作 buffer 覆盖导致的 down 梯度全零已修复。 +- gate/up/down 三组权重均有实际更新证据。 +- AMX BF16/NUMA subpool 已替代极慢的标量单核基座梯度循环。 + +### 6.2 尚未完成 + +- `1e16~1e17` 梯度量级仍需分阶段定位;`PASS` 只表示结构检查通过。 +- 需要 PyTorch 小尺寸参考梯度逐元素验证 gate/up/down,不能只检查非零。 +- 需要 TP=1、TP=2,且包含 `expert_id>0` 的 stride 覆盖测试。 +- 需要连续两步激活不同 expert,验证 inactive expert 无残留梯度。 +- 需要 `full/hybrid/lora` 三模式回归,确认纯 LoRA 行为不变。 +- 个人分支仍需合入共同基线之后的官方 main 提交并解决潜在冲突。 + +### 6.3 容易误读的日志 + +- Trainer 的 `grad_norm=0` 可能只统计 model tree 中的 named parameters;KT 注入 optimizer + 的 expert Parameter 可能不在该统计路径中。 +- `Number of trainable params = 0` 也可能来自 HF expert zero-storage placeholder;应检查 + optimizer 中是否存在 144 个 KT expert Parameter。 +- requant 次数只能证明调用发生,不能证明权重变化。 +- loss 下降可能来自 attention/router 等非 expert 参数。 +- 自动 summary 可能漏报 child SIGSEGV;崩溃以原始 `train.log`、GDB 和 child exitcode 为准。 +- 一步测试若 LR=0,`changed=0` 不是更新链路失败证据。 + +## 7. 开发与验证要求 + +### 7.1 修改前 + +1. 运行 `git status --short --branch`,不要覆盖现有改动。 +2. 记录目标模式、TP 数、NUMA 数、`E/H/F/I`、dtype 和 `lora_rank`。 +3. 使用符号名定位,不依赖本文记录的历史行号。 +4. 判断修改属于 Python Full-FT 接入、顶层 TP 调度还是 AMX 子核,不要跨层打补丁。 + +### 7.2 构建 + +普通性能构建使用 Release;需要 GDB 源码行时使用同一 Python 环境构建 +`RelWithDebInfo`: ```bash cd /mnt/data2/wbw/ktransformers @@ -190,62 +469,77 @@ CPUINFER_BUILD_TYPE=RelWithDebInfo \ --no-build-isolation --no-deps --force-reinstall ./kt-kernel ``` -确认测试实际 import 的 `.so` 含 `.debug_info`/`.debug_line`。然后使用 expert-only -runner 的 `--gdb`;只有磁盘空间足够时才使用 `--gdb-core`。修复后的 expert-only -短训练仍是必需验收项。 - -1. **静态检查**:确认 setter 不在 LoRA 条件分支内;确认所有指针先判空再偏移。 -2. **构建**:重新编译并安装 `kt_kernel_ext`。 -3. **参考梯度测试**: - - 分别覆盖 TP=1 和 TP=2。 - - 使用 `expert_id > 0`,避免错误 expert stride 被 expert 0 掩盖。 - - 与 PyTorch outer-product 参考值比较 gate/up/down 三组梯度。 - - TP=2 时分别检查两个 `F` 分片,确认无覆盖。 -4. **跨 step 测试**:连续两步激活不同 expert,确认未激活 expert 不保留旧梯度。 -5. **模式回归**:`full_weight_grad=false` 不写基座梯度;纯 LoRA 结果不变。 -6. **短训练**: - - step 内三个 `grad_*_proj_buf.abs().max() > 0`。 - - optimizer 后至少一个 `*_proj_buf` 发生有限变化。 - - requant 后下一次 forward 使用更新后的权重。 - -不要用 loss 下降证明 Full FT 生效;attention、router 或其他非 expert 参数也能 -让 loss 下降。 +安装后必须确认测试实际 import 的 `.so` 与仓库 build 产物一致,并检查 debug build +包含 `.debug_info`/`.debug_line`。 + +### 7.3 验证顺序 + +1. 静态检查:setter 不在 LoRA 条件分支内;所有指针先判空再 offset。 +2. 格式与构建:`clang-format --dry-run --Werror`,AMX/CUDA build 成功。 +3. AMX tile 单测:BF16 `32x32xK` 与标量 BF16 参考比较。 +4. 小尺寸参考梯度:gate/up/down 分别与 PyTorch outer-product 比较。 +5. TP 测试:TP=1、TP=2,检查两个 `F` slice 无覆盖、无缺口。 +6. 跨 step 测试:不同 active expert,检查清零和残留。 +7. 模式回归:`full_weight_grad=false` 不写基座梯度;LoRA 结果不变。 +8. 短训练:三个 grad buffer 非零且有限,optimizer 后权重发生有限变化,下一次 + forward 使用 requant 后的新权重。 +9. 性能测试:关闭 GDB 和重型 probe,至少 15 step,前 5 step warmup,后 10 step TPS。 ## 8. 完成标准 -- C++ 子核实际收到 `full_weight_grad=true`。 -- 三组基座梯度非零、有限,并与参考实现一致。 -- TP 分片覆盖完整 `F`,无重叠、越界或错误 stride。 -- optimizer 能找到三个基座 Parameter,并在 step 后更新它们。 -- inactive expert 无跨 step 残留梯度。 -- `full`、`hybrid`、`lora` 三种模式行为符合各自契约。 -- 纯 LoRA 和现有 AMX MoE 测试无回归。 +结构正确性: + +- 所有子核收到正确的 Full-FT 开关和 TP slice。 +- gate/up/down C++ grad 与 `Parameter.grad` 一致。 +- optimizer 更新权威 BF16 Parameter,requant 使用更新后的指针。 +- 无 SIGSEGV、越界、TP 覆盖或跨 step 残留。 + +数值正确性: + +- 三组梯度与 PyTorch 参考实现误差在约定容差内。 +- 梯度尺度有合理解释,不仅仅是“finite/nonzero”。 +- full/hybrid/lora 均无回归。 + +性能正确性: -## 9. 调试资料 +- TPS 结果不包含 GDB 和全量梯度扫描开销。 +- 报告明确 batch size、sequence length、GAS、warmup 和稳定 step 数。 +- CPU 内存、GPU 显存、backward、optimizer、requant 分项可追溯。 + +只有同时满足结构、数值和模式回归,才可宣称 Full FT 完全正确。目前已确认结构链路和 +权重更新生效,但梯度尺度问题仍未关闭。 + +## 9. 调试资料索引 ```text FFTtest/Qwen3-30B-A3B/test_log/20260710_103230_1gpu_AMX_BF16_FULLFT_EXPERTONLY/ - summary.md - expert_weight_change_check.{txt,json} - phase4/expert_buf_probe.json - phase4/train.log + # 15 步但 changed=0/12;证明初始 Full-FT 链路未生效 FFTtest/Qwen3-30B-A3B/test_log/20260710_192456_1gpu_AMX_BF16_FULLFT_EXPERTONLY/ - phase4/train.log # exitcode=-11,Signal 11 (SIGSEGV) - phase4/log_analysis.txt # 正确检测到 SIGSEGV - phase4/exit_code.txt # 外层 launcher 退出码 1 - summary.md # P5 旧结论错误,不得作为崩溃判据 + phase4/train.log # child exitcode=-11 / Signal 11 + phase4/log_analysis.txt + summary.md # 历史汇总漏报崩溃,不得单独采用 FFTtest/Qwen3-30B-A3B/test_log/20260713_105328_1gpu_AMX_BF16_FULLFT_EXPERTONLY/ - phase4/gdb_sigsegv.log # t=8 越界、非法 tok_pos 和双 NUMA 原生栈 - phase4/train.log - summary.md + phase4/gdb_sigsegv.log # 路由缓存布局误用的直接证据 + +FFTtest/Qwen3-30B-A3B/test_log/20260713_111729_1gpu_AMX_BF16_FULLFT_EXPERTONLY/ + # 1 step、LR=0;只证明无 SIGSEGV,不能判断权重更新 + +FFTtest/Qwen3-30B-A3B/test_log/20260713_120636_1gpu_AMX_BF16_FULLFT_EXPERTONLY/ + expert_gradient_check.txt # gate/up 非零,down 0/48 + expert_weight_change_check.txt + phase4/step_timing/ + +FFTtest/Qwen3-30B-A3B/test_log/20260713_140101_1gpu_AMX_BF16_FULLFT_EXPERTONLY/ + expert_gradient_check.txt # gate/up/down 144/144 + expert_weight_change_check.txt # changed=9/12,down=4/4 + phase4/step_timing/ # backward 约 22.8s,probe 进入 other -FFTtest/Qwen3-30B-A3B/expert_buf_probe.py FFTtest/Qwen3-30B-A3B/gdb_sigsegv.gdb FFTtest/Qwen3-30B-A3B/run_full_ft_test_1gpu_bf16_frozen.sh --gdb +FFTtest/Qwen3-30B-A3B/expert_buf_probe.py ``` -该 run 的直接证据是 `changed=0/12`、`max_abs_delta=0`;训练 loss -仍下降不能否定 expert 基座梯度失效。后续 run 的直接崩溃证据必须以原始 -`phase4/train.log` 和 `phase4/gdb_sigsegv.log` 为准,不能只依赖生成的 summary。 +调试时优先读取原始 `phase4/train.log`、GDB backtrace、逐步 timing 和 probe JSON;自动 +生成的 `summary.md` 只能作为索引,不能覆盖原始证据。 From 06f06065697b1e50ba6eefc10ab85b3cf37e61a6 Mon Sep 17 00:00:00 2001 From: illu Date: Mon, 13 Jul 2026 07:55:02 +0000 Subject: [PATCH 06/20] [docs](kt-kernel): align fork remote terminology --- kt-kernel/AGENTS.md | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/kt-kernel/AGENTS.md b/kt-kernel/AGENTS.md index c404b32f4..691e27e69 100644 --- a/kt-kernel/AGENTS.md +++ b/kt-kernel/AGENTS.md @@ -14,10 +14,11 @@ embedding 等参数的更新当作 expert Full FT 已生效的证据。 截至 2026-07-13,本地比较对象为: -- 官方仓库:`kvcache-ai/ktransformers`,remote 名为 `origin`。 +- 官方仓库:`kvcache-ai/ktransformers`,remote 名为 `upstream`。 - 个人仓库:`Illumination111/ktransformers-fullFT_development`,remote 名为 - `development`。 -- 个人开发分支:`fullft-development`,当前提交 `5ce0767`。 + `origin`。 +- 个人开发分支:`fullft-development`;Full-FT 核心生产代码截止 `5ce0767`, + 后续纯文档提交不改变该代码基线。 - 共同基线:官方提交 `8e46e58`。 - Full FT 初始快照:`e99b5e1`,来源包含 poryfly 的 `239bac5`、`25fa2fd`。 - 后续修复:`2d81e86`、`20f645c`、`5ce0767`。 From 01c0c929e01705c767c292d48338d025035ac732 Mon Sep 17 00:00:00 2001 From: illu Date: Wed, 15 Jul 2026 05:40:31 +0000 Subject: [PATCH 07/20] [chore](kt-kernel): keep agent notes local --- kt-kernel/AGENTS.md | 546 -------------------------------------------- 1 file changed, 546 deletions(-) delete mode 100644 kt-kernel/AGENTS.md diff --git a/kt-kernel/AGENTS.md b/kt-kernel/AGENTS.md deleted file mode 100644 index 691e27e69..000000000 --- a/kt-kernel/AGENTS.md +++ /dev/null @@ -1,546 +0,0 @@ -# AGENTS.md:kt-kernel MoE Full FT 开发、差异与调试记录 - -本文件作用于 `kt-kernel/` 及其全部子目录,用于说明个人 Full Fine-Tuning -(Full FT)分支相对官方仓库的真实代码差异,并保留从“专家权重不更新”到 -“SIGSEGV、down 梯度为零、性能过低”这一系列问题的试错记录。 - -这里的 Full FT 专指 KTransformers CPU/AMX MoE expert 基座权重 -`gate_proj`、`up_proj`、`down_proj` 的训练,不把普通 GPU attention、router、 -embedding 等参数的更新当作 expert Full FT 已生效的证据。 - -## 1. 仓库、基线与比较口径 - -### 1.1 分支关系 - -截至 2026-07-13,本地比较对象为: - -- 官方仓库:`kvcache-ai/ktransformers`,remote 名为 `upstream`。 -- 个人仓库:`Illumination111/ktransformers-fullFT_development`,remote 名为 - `origin`。 -- 个人开发分支:`fullft-development`;Full-FT 核心生产代码截止 `5ce0767`, - 后续纯文档提交不改变该代码基线。 -- 共同基线:官方提交 `8e46e58`。 -- Full FT 初始快照:`e99b5e1`,来源包含 poryfly 的 `239bac5`、`25fa2fd`。 -- 后续修复:`2d81e86`、`20f645c`、`5ce0767`。 - -官方 `origin/main` 当时已前进到 `7c021b4`,比共同基线多 3 个与本轮 Full FT -基本无关的提交: - -```text -79b265b normalize compressed RAWINT4 weights -cb9f47d detect bound ports before launch -7c021b4 bump sglang submodule -``` - -因此,统计“个人做了什么”时必须使用共同基线或三点 diff: - -```bash -git merge-base origin/main fullft-development -git diff --shortstat origin/main...fullft-development -git diff --numstat origin/main...fullft-development -``` - -不要直接把 `git diff origin/main fullft-development` 的全部结果都归到个人修改; -那样会把个人分支尚未合入的 3 个官方提交也统计为差异。 - -### 1.2 修改规模 - -个人 Full FT 分支相对共同基线 `8e46e58`: - -| 口径 | 文件数 | 新增 | 删除 | 说明 | -|---|---:|---:|---:|---| -| 当前工作树(含本次扩写后的文档) | 16 | 1581 | 160 | 全部位于 `kt-kernel/` | -| 仅生产代码 | 15 | 1036 | 160 | 总代码 churn 约 1196 行 | -| Python SFT 接入 | 11 | 658 | 143 | 约占生产代码改动的三分之二 | -| C++ config/binding/AMX/TP | 4 | 378 | 17 | 约占生产代码改动的三分之一 | - -应区分两个阶段: - -1. `e99b5e1` 是完整 Full FT 功能接入,涉及 15 个代码文件,`+809/-155`, - 修改面较广。 -2. 后续 3 个 debug/fix 提交的生产代码只修改两个文件,合计约 `+255/-33`: - `operators/amx/sft_moe.hpp` 和 `operators/moe-sft-tp.hpp`。 - -提交 `5ce0767` 时旧版本文档为 251 行;本次仅扩写说明和试错记录,没有改变上述 -15 个生产代码文件的统计。 - -因此,“个人 Full FT 分支总体改动”属于中等偏大;“本轮 SIGSEGV/down 梯度/ -AMX 性能 debug”则高度集中在两个 `.hpp` 文件。 - -## 2. 官方代码与个人 Full FT 代码的行为差异 - -### 2.1 官方基线的限制 - -官方基线的 SFT expert 路径以 LoRA 为主要训练对象: - -- 没有 `full/hybrid/lora` 的明确 KT train-mode 映射。 -- `lora_rank` 必须为正数,纯 Full FT 的 `lora_rank=0` 不成立。 -- expert 原始权重只用于加载和 AMX forward,不是 optimizer 可见的 - `nn.Parameter`。 -- Python autograd 不把三个 expert 基座权重作为输入,也不返回对应梯度。 -- C++ `MOESFTConfig` 没有 `full_weight_grad` 和三个基座梯度指针。 -- C++ backward 不会把 gate/up/down 基座梯度写回 Python tensor。 -- optimizer 更新后没有把 BF16 expert 基座权重重新量化到 AMX BufferB 的闭环。 - -所以,仅在 LLaMA-Factory 中设置 `finetuning_type: full`,不能证明 KT CPU experts -正在 Full FT;attention/router 等 GPU 参数仍可能让 loss 下降。 - -### 2.2 个人分支建立的完整数据链路 - -```text -LLaMA-Factory finetuning_type - -> KTConfig.kt_train_mode / kt_full_weight_grad - -> wrapper 创建 CPU BF16 gate/up/down nn.Parameter - -> KTMoEFunction.forward 显式接收三个 Parameter - -> C++ backward 写 grad_gate/up/down_proj_buf - -> KTMoEFunction.backward 返回三个基座梯度 - -> Trainer optimizer 注入并更新三个 *_proj_buf - -> optimizer.step 后标记 _base_weights_dirty - -> 下一次 forward 调用 update_base_weights - -> BF16 基座权重重新量化到 AMX BufferB -``` - -权威数据定义: - -- `gate_proj_buf`、`up_proj_buf`、`down_proj_buf`:optimizer 可见的 CPU BF16 - expert 基座权重。 -- `grad_gate_proj_buf`、`grad_up_proj_buf`、`grad_down_proj_buf`:C++ backward - 写入、autograd 返回的基座梯度。 -- AMX 量化权重:forward 使用的派生副本,optimizer 后必须重建。 -- HF model tree 中的 expert 权重:清理为 zero-storage placeholder,不再是 Full FT - 的权威副本,也不能用于训练前后权重比较。 - -### 2.3 模式语义 - -- `full`:训练 expert 基座权重,通常 `lora_rank=0`。 -- `hybrid`:同时训练 expert 基座权重和 LoRA,允许 `lora_rank>0`。 -- `lora`:只训练 LoRA,不得计算、同步或更新 expert 基座权重。 - -修改任何共享路径前必须分别确认三种模式,不能用 Full FT 修复破坏 LoRA。 - -## 3. 逐文件说明:个人分支具体修改了什么 - -路径均相对 `kt-kernel/`。 - -### 3.1 Python 训练接入层 - -| 文件 | 相对官方的主要修改 | -|---|---| -| `python/sft/config.py` | 增加 `kt_train_mode`、`kt_full_weight_grad`;读取 `ACCELERATE_KT_TRAIN_MODE`,将 `full/hybrid` 映射为基座梯度开启。 | -| `python/experts.py` | `KTMoEWrapper` 和 SFT wrapper factory 接受并传递 `full_weight_grad`。 | -| `python/sft/wrapper.py` | 将 LLaMA-Factory `finetuning_type` 映射到 KT train mode;把 `lora_rank=0` 当作合法纯 Full FT;初始化三个权威 BF16 Parameter/grad buffer;释放重复权重;把 HF expert 权重替换为 zero-storage placeholder。 | -| `python/sft/base.py` | 保存 `_full_weight_grad` 和 `_base_weights_dirty`;创建 `[E,F,H]` gate/up 与 `[E,H,F]` down Parameter/grad buffer;允许 Full FT 在无 LoRA 时运行;定义 `update_base_weights()` 接口;避免 `lora_rank=0` 除零。 | -| `python/sft/layer.py` | Full FT 时强制进入 autograd 路径;将三个基座 Parameter 传入 `KTMoEFunction`;optimizer 后发现 dirty 权重时触发 requant;Full FT 时保留 router 梯度;兼容 transformers v5 TopKRouter/GLM4 路由。 | -| `python/sft/autograd.py` | forward 增加三个 expert 基座 Parameter 输入;backward 返回 C++ 写入的 gate/up/down 梯度,使 PyTorch 给 Parameter 建立 `.grad`。 | -| `python/sft/lora.py` | 拆分 LoRA 参数与 Full-FT 基座参数收集;把 48 层 × 3 投影共 144 个 Parameter 注入 optimizer;纯 Full FT 跳过 LoRA buffer;分布式时同步基座梯度;optimizer 后标记基座权重 dirty。 | -| `python/sft/amx.py` | 将 Full-FT 开关和三个梯度 data pointer 写入 C++ config/backward task;`lora_rank=0` 时传空 LoRA 指针;增加 `update_base_weights()`,优先复用 C++ 对象并重新量化,缺少 binding 时才完整重建。 | -| `python/sft/weights.py` | 明确 `*_proj_buf` 为 Full FT 权威权重;清理 model tree 中的冗余 expert 参数并打 `_kt_zero_storage` 标记,避免重复计数和重复占内存。 | -| `python/sft/arch.py` | 增加 GLM4 MoE 架构识别;这属于同期兼容性修改,不是本次 Qwen3 Full FT bugfix 的核心。 | -| `python/sft/__init__.py` | 导出新增的 Full-FT 参数收集和相关 SFT API。 | - -### 3.2 C++ 配置与绑定 - -| 文件 | 相对官方的主要修改 | -|---|---| -| `operators/common.hpp` | 在 `MOESFTConfig` 中增加 `full_weight_grad` 和 `grad_gate/up/down_proj` 三个零拷贝指针。 | -| `ext_bindings.cpp` | backward binding 接收并转发三个基座梯度指针;向 Python 暴露 Full-FT config 字段;增加 `set_base_weight_pointers()`,支持 optimizer 后复用既有 C++ MoE 对象并 requant。 | - -### 3.3 两个核心 `.hpp` - -#### `operators/moe-sft-tp.hpp` - -该文件负责顶层 TP 调度和完整梯度 tensor 的切片: - -1. 子核创建后显式调用 `set_full_weight_grad()`,补回派生 config 被 - `GeneralMOEConfig` slicing 丢失的字段。 -2. 传播逻辑移到 `if constexpr (!kSkipLoRA)` 外,确保纯 Full FT 的 - `lora_rank=0` 仍启用基座梯度。 -3. backward dispatch 前按 TP offset 计算 gate/up/down 三个 slice 指针。 -4. 将完整 `F` 传给子核作为 global stride,而不是误用当前 TP 的本地 `I`。 -5. 检查所有 TP slice 对完整 `F` 的覆盖,无重叠、无缺口。 -6. dispatch 前统一并行清零三组完整梯度,避免跨 step 残留和多个 TP 子核竞争清零。 -7. 空指针必须先判断再做 offset,避免对 `nullptr` 做未定义指针运算。 - -#### `operators/amx/sft_moe.hpp` - -该文件负责每个 NUMA/TP 子核的真实 backward 和基座梯度计算: - -1. 保存并更新 `sft_config_.full_weight_grad`。 -2. `backward_base_weight_grad` 同时接收本地 `I` 和完整 `F`;本地计算用 `I`, - 写回 expert/row stride 用 `F`。 -3. gate/up/down 三组指针全部有效时才计算基座梯度。 -4. 扩大/复用工作区,保证基座 FP32 累加需要的容量。 -5. `lora_rank=0` 时 scaling 为 0,完成 gate/up grad-input 后跳过 LoRA remainder, - 不进入零秩临时区路径。 -6. 不再把 token-major 路由表当作 expert-major token list;直接读取 backward 已 - 打包好的 `m_local_input_ptr_[expert]` 和 expert-major grad-output。 -7. 为 down 梯度保存独立、只读、已经乘过 router weight 的 `grad_output` 快照; - gate/up grad-input 可以继续复用工作 buffer,但不能覆盖 down 所需的 dY。 -8. 使用 AMX BF16 `32x32` output tile、FP32 累加计算基座梯度。 -9. 任务按 expert、projection、tile 投递给当前 NUMA subpool,而不是在每个 NUMA - 节点只用单核执行三重标量循环。 -10. gate/up 共用输入 tile,每个任务写独占输出区域,避免锁和写竞争。 - -暂不修改通用 `operators/moe-tp.hpp`。将其改成保留派生 config 的泛型重构会影响 -所有 MoE backend,应另开改动并做完整回归。 - -## 4. TP 梯度布局契约 - -定义: - -- `E`:expert 数量。 -- `H`:hidden size。 -- `F`:完整 intermediate size。 -- `I`:当前 TP 子核的本地 intermediate size。 -- `tp_offset`:当前 TP slice 在完整 `F` 维的起点。 - -完整梯度布局: - -```text -gate/up: [E, F, H] -down: [E, H, F] -``` - -顶层 TP wrapper 传给子核的起始指针: - -```text -gate/up = base + tp_offset * H -down = base + tp_offset -``` - -子核写回公式: - -```text -gate/up: expert * F * H + local_i * H + h -down: expert * H * F + h * F + local_i -``` - -易错点: - -- 计算循环范围是本地 `I`,但 expert/row stride 必须是完整 `F`。 -- 所有 TP 必须写入互不重叠的区域。 -- `expert_id=0` 可能掩盖错误 expert stride,参考测试必须包含 `expert_id>0`。 -- 梯度清零只能由顶层统一完成,不能让多个 TP 子核同时清相同完整 buffer。 - -## 5. 修改与试错时间线 - -本节保留失败路径,因为这些失败揭示了仅看 loss、requant 日志或单步权重抽检会得出 -错误结论。 - -### 5.1 阶段 0:官方路径不能形成 expert Full FT 闭环 - -最初仅设置 LLaMA-Factory `finetuning_type=full`,但 KT expert 权重不是 optimizer -可见 Parameter,C++ 也没有基座梯度输出。loss 变化最多说明其他参数在训练,不能说明 -CPU experts 更新。 - -结论:必须建立 Parameter → C++ grad → autograd → optimizer → requant 的完整闭环。 - -### 5.2 阶段 1:`e99b5e1` 接入 Full FT,但专家权重仍不更新 - -`e99b5e1` 加入 Python/C++ Full-FT 数据链路。随后运行: - -```text -20260710_103230_1gpu_AMX_BF16_FULLFT_EXPERTONLY -``` - -测试已冻结非 expert 参数,并向 optimizer 注入 144 个 expert 基座 Parameter,但结果为: - -- 15 步 `grad_norm=0`。 -- step 0 到 step 15 抽检 `changed=0/12`。 -- `max_abs_delta=0`。 -- 有 720 次 requant 日志。 - -试错结论: - -- optimizer 参数数量正确,不等于 C++ 真正产生了梯度。 -- requant 被调用,不等于其输入权重发生了变化。 -- 问题继续下沉到 C++ 子核配置和 TP 梯度写回。 - -根因之一是顶层 `TP_MOE_SFT` 收到 `full_weight_grad=true` 后,创建子核时把派生配置 -赋给 `GeneralMOEConfig`,发生 slicing;子核重建 `MOESFTConfig` 后开关恢复为 false, -`backward_base_weight_grad` 被静默跳过。 - -### 5.3 阶段 2:`2d81e86` 修复开关/TP stride 后出现 SIGSEGV - -`2d81e86` 尝试修复: - -- 将 `full_weight_grad` 显式传播到子核。 -- 给每个 TP 传独立梯度 slice。 -- 使用完整 `F` 作为全局 stride。 -- 扩大临时区并统一清零。 - -但测试: - -```text -20260710_192456_1gpu_AMX_BF16_FULLFT_EXPERTONLY -``` - -在第一个训练 step 前后发生 SIGSEGV: - -- child `exitcode=-11`,原始日志明确记录 `Signal 11 (SIGSEGV)`。 -- 外层 `exit_code.txt=1` 只是 accelerate launcher 退出码。 -- 当时的自动 `summary.md` 错误写成“未检测到崩溃”,不能作为反证。 -- 该 run 没有 C++ backtrace,最初只能假设是 gradient index/stride 越界。 - -失败原因:修复了“是否计算”和“写到哪里”,但新实现为了计算基座梯度又错误理解了 -路由缓存布局。 - -### 5.4 阶段 3:加入带符号 GDB,定位错误路由表解释 - -为避免继续猜测,在 FFTtest runner 中增加 batch GDB,并用同一 Kllama Python 环境 -构建 `RelWithDebInfo` 扩展。日志: - -```text -20260713_105328_1gpu_AMX_BF16_FULLFT_EXPERTONLY -``` - -GDB 证据: - -- `m_local_pos_cache` 的真实布局是 `[token_idx][route_slot]`。 -- 每个 token 的 route 长度是 `k=8`。 -- 旧代码却按 `m_local_pos_cache[expert_idx][t]` 把它当作 expert token list。 -- 第一个 expert 执行到 `t=8` 即越界,读出垃圾 `tok_pos=81349952`。 -- 最终在读取 `input_row[0]` 时于当时的 `sft_moe.hpp:1968` 触发 SIGSEGV。 -- 两个 NUMA 子核都进入了相同错误路径。 - -这一步推翻了“只是 TP output stride 错误”的单一假设。崩溃发生在读取输入行,而不是 -最终写回梯度的位置。 - -### 5.5 阶段 4:`20f645c` 最小修复消除 SIGSEGV,但单步测试无效 - -`20f645c` 不再反查 token-major 路由表,改用 backward 已生成的 expert-major packed -buffer: - -```text -m_local_input_ptr_[expert_idx] -grad_output_bf16_ptr_[expert_idx] -``` - -随后单步测试: - -```text -20260713_111729_1gpu_AMX_BF16_FULLFT_EXPERTONLY -``` - -结果: - -- GDB 显示进程正常退出,无 SIGSEGV。 -- backward 约 890.9 秒。 -- 权重抽检仍是 `changed=0/12`。 - -这里不能得出“梯度仍无效”的结论,因为只运行 1 step,而第一个 optimizer step 的 -learning rate 是 0。这个 run 证明了崩溃消失,但不能验证权重更新。 - -试错教训:至少需要 2 个 optimizer step,最好 3~5 步;必须同时记录每步 LR。 - -### 5.6 阶段 5:三步测试证明 gate/up 更新,但暴露 down 梯度为零 - -继续运行: - -```text -20260713_120636_1gpu_AMX_BF16_FULLFT_EXPERTONLY -``` - -结果: - -- 3 个 step 正常退出,无 SIGSEGV。 -- 权重抽检 `changed=5/12`,证明 Full-FT optimizer 链路总体已经能更新 expert。 -- gate 梯度 `48/48` 非零。 -- up 梯度 `48/48` 非零。 -- down 梯度 `0/48`,所以总计只有 `96/144` 非零。 -- backward 平均约 754.5 秒/step。 -- 梯度最大值达到 `4e16~9e16`,虽有限但数值明显异常。 - -为了区分“C++ buffer 有值但 autograd/optimizer 看不到”和“C++ 本身没有写值”,测试脚本 -增加了对 `grad_*_proj_buf` 与独立 `Parameter.grad` 的全量扫描。二者的非零数量和最大值 -一致,说明 down=0 发生在 C++ 计算链路,而不是 Python autograd 丢梯度。 - -进一步检查发现,down 基座梯度需要的 route-weighted dY 原本位于 -`grad_output_bf16_ptr_`,但 gate/up grad-input backward 随后复用了同一 buffer,down 计算 -时读到的内容已被覆盖。 - -失败方案:简单地把 `grad_output_bf16_ptr_` 当作长期只读 down dY。该指针只是工作 buffer, -生命周期不满足要求。 - -### 5.7 阶段 6:`5ce0767` 保存 down 快照并用 AMX/NUMA 并行 - -`5ce0767` 同时处理正确性和性能: - -- 在工作 buffer 被 gate/up 路径复用前,保存 route-weighted dY 的独立只读快照。 -- down 基座梯度只读取该快照。 -- gate/up/down 基座梯度改用 AMX BF16 `32x32` tile、FP32 累加。 -- 任务提交给每个 NUMA 节点已有的 subpool,而不是每节点单核标量循环。 -- 纯 Full FT `lora_rank=0` 跳过 LoRA remainder。 - -验证日志: - -```text -20260713_140101_1gpu_AMX_BF16_FULLFT_EXPERTONLY -``` - -结果: - -- 5/5 step 正常退出,无 SIGSEGV。 -- C++ grad buffer `144/144` 非零,`Parameter.grad` 也是 `144/144` 非零。 -- gate/up/down 分别为 `48/48`、`48/48`、`48/48`。 -- 权重抽检 `changed=9/12`;其中 down 抽检 `4/4` 全部变化。 -- 最大权重差约 `3.8147e-05`,没有非有限权重。 -- backward 从约 754.5 秒降到约 22.8 秒,约 33 倍加速、97% 降幅。 - -但该 run 只证明“链路连通、无非有限值、权重可更新”,不能证明数值完全正确: - -- 梯度最大值仍在 `1e16~1e17`。 -- loss 约 1~2 时,这个量级明显可疑。 -- AdamW 会归一化梯度,权重变化合理不能反推原始梯度合理。 - -### 5.8 探针本身造成的性能误判 - -全量梯度探针每个 step 扫描: - -- 144 个 C++ grad buffer。 -- 144 个独立 `Parameter.grad`。 -- 合计约 580 亿个 BF16 元素。 -- `isfinite`、`count_nonzero`、`abs().max()` 会重复遍历。 - -在 `20260713_140101` 中: - -- backward 已降到约 22.8 秒/step。 -- `step_other` 仍约 256.7 秒/step,占 61.8%。 -- 该 `other` 主要是 pre-optimizer 全量梯度扫描,不是 GDB。 - -因此诊断 run 与性能 run 必须分开: - -- 数值诊断:1~3 step,抽样或分阶段记录 C++ 中间量。 -- 性能测试:不使用 GDB、不扫描 expert 梯度/权重,只保留内存、显存和轻量 step timing。 -- 正式 TPS:固定 15 step,跳过前 5 个 warmup,用后 10 个 step 总 token/总时间计算。 - -## 6. 当前状态与仍未解决的问题 - -### 6.1 已确认解决 - -- `full_weight_grad` 能从 Python 到达顶层和所有 TP/NUMA 子核。 -- 三个 expert 基座 Parameter 能进入 optimizer。 -- TP gradient slice、global stride 和逐步清零已修复。 -- 错误解释 `m_local_pos_cache` 导致的 SIGSEGV 已修复。 -- down dY 被工作 buffer 覆盖导致的 down 梯度全零已修复。 -- gate/up/down 三组权重均有实际更新证据。 -- AMX BF16/NUMA subpool 已替代极慢的标量单核基座梯度循环。 - -### 6.2 尚未完成 - -- `1e16~1e17` 梯度量级仍需分阶段定位;`PASS` 只表示结构检查通过。 -- 需要 PyTorch 小尺寸参考梯度逐元素验证 gate/up/down,不能只检查非零。 -- 需要 TP=1、TP=2,且包含 `expert_id>0` 的 stride 覆盖测试。 -- 需要连续两步激活不同 expert,验证 inactive expert 无残留梯度。 -- 需要 `full/hybrid/lora` 三模式回归,确认纯 LoRA 行为不变。 -- 个人分支仍需合入共同基线之后的官方 main 提交并解决潜在冲突。 - -### 6.3 容易误读的日志 - -- Trainer 的 `grad_norm=0` 可能只统计 model tree 中的 named parameters;KT 注入 optimizer - 的 expert Parameter 可能不在该统计路径中。 -- `Number of trainable params = 0` 也可能来自 HF expert zero-storage placeholder;应检查 - optimizer 中是否存在 144 个 KT expert Parameter。 -- requant 次数只能证明调用发生,不能证明权重变化。 -- loss 下降可能来自 attention/router 等非 expert 参数。 -- 自动 summary 可能漏报 child SIGSEGV;崩溃以原始 `train.log`、GDB 和 child exitcode 为准。 -- 一步测试若 LR=0,`changed=0` 不是更新链路失败证据。 - -## 7. 开发与验证要求 - -### 7.1 修改前 - -1. 运行 `git status --short --branch`,不要覆盖现有改动。 -2. 记录目标模式、TP 数、NUMA 数、`E/H/F/I`、dtype 和 `lora_rank`。 -3. 使用符号名定位,不依赖本文记录的历史行号。 -4. 判断修改属于 Python Full-FT 接入、顶层 TP 调度还是 AMX 子核,不要跨层打补丁。 - -### 7.2 构建 - -普通性能构建使用 Release;需要 GDB 源码行时使用同一 Python 环境构建 -`RelWithDebInfo`: - -```bash -cd /mnt/data2/wbw/ktransformers -CPUINFER_BUILD_TYPE=RelWithDebInfo \ - /mnt/data2/wbw/conda/envs/Kllama/bin/python3.12 \ - kt-kernel/setup.py build_ext --inplace -CPUINFER_BUILD_TYPE=RelWithDebInfo \ - /mnt/data2/wbw/conda/envs/Kllama/bin/python3.12 -m pip install \ - --no-build-isolation --no-deps --force-reinstall ./kt-kernel -``` - -安装后必须确认测试实际 import 的 `.so` 与仓库 build 产物一致,并检查 debug build -包含 `.debug_info`/`.debug_line`。 - -### 7.3 验证顺序 - -1. 静态检查:setter 不在 LoRA 条件分支内;所有指针先判空再 offset。 -2. 格式与构建:`clang-format --dry-run --Werror`,AMX/CUDA build 成功。 -3. AMX tile 单测:BF16 `32x32xK` 与标量 BF16 参考比较。 -4. 小尺寸参考梯度:gate/up/down 分别与 PyTorch outer-product 比较。 -5. TP 测试:TP=1、TP=2,检查两个 `F` slice 无覆盖、无缺口。 -6. 跨 step 测试:不同 active expert,检查清零和残留。 -7. 模式回归:`full_weight_grad=false` 不写基座梯度;LoRA 结果不变。 -8. 短训练:三个 grad buffer 非零且有限,optimizer 后权重发生有限变化,下一次 - forward 使用 requant 后的新权重。 -9. 性能测试:关闭 GDB 和重型 probe,至少 15 step,前 5 step warmup,后 10 step TPS。 - -## 8. 完成标准 - -结构正确性: - -- 所有子核收到正确的 Full-FT 开关和 TP slice。 -- gate/up/down C++ grad 与 `Parameter.grad` 一致。 -- optimizer 更新权威 BF16 Parameter,requant 使用更新后的指针。 -- 无 SIGSEGV、越界、TP 覆盖或跨 step 残留。 - -数值正确性: - -- 三组梯度与 PyTorch 参考实现误差在约定容差内。 -- 梯度尺度有合理解释,不仅仅是“finite/nonzero”。 -- full/hybrid/lora 均无回归。 - -性能正确性: - -- TPS 结果不包含 GDB 和全量梯度扫描开销。 -- 报告明确 batch size、sequence length、GAS、warmup 和稳定 step 数。 -- CPU 内存、GPU 显存、backward、optimizer、requant 分项可追溯。 - -只有同时满足结构、数值和模式回归,才可宣称 Full FT 完全正确。目前已确认结构链路和 -权重更新生效,但梯度尺度问题仍未关闭。 - -## 9. 调试资料索引 - -```text -FFTtest/Qwen3-30B-A3B/test_log/20260710_103230_1gpu_AMX_BF16_FULLFT_EXPERTONLY/ - # 15 步但 changed=0/12;证明初始 Full-FT 链路未生效 - -FFTtest/Qwen3-30B-A3B/test_log/20260710_192456_1gpu_AMX_BF16_FULLFT_EXPERTONLY/ - phase4/train.log # child exitcode=-11 / Signal 11 - phase4/log_analysis.txt - summary.md # 历史汇总漏报崩溃,不得单独采用 - -FFTtest/Qwen3-30B-A3B/test_log/20260713_105328_1gpu_AMX_BF16_FULLFT_EXPERTONLY/ - phase4/gdb_sigsegv.log # 路由缓存布局误用的直接证据 - -FFTtest/Qwen3-30B-A3B/test_log/20260713_111729_1gpu_AMX_BF16_FULLFT_EXPERTONLY/ - # 1 step、LR=0;只证明无 SIGSEGV,不能判断权重更新 - -FFTtest/Qwen3-30B-A3B/test_log/20260713_120636_1gpu_AMX_BF16_FULLFT_EXPERTONLY/ - expert_gradient_check.txt # gate/up 非零,down 0/48 - expert_weight_change_check.txt - phase4/step_timing/ - -FFTtest/Qwen3-30B-A3B/test_log/20260713_140101_1gpu_AMX_BF16_FULLFT_EXPERTONLY/ - expert_gradient_check.txt # gate/up/down 144/144 - expert_weight_change_check.txt # changed=9/12,down=4/4 - phase4/step_timing/ # backward 约 22.8s,probe 进入 other - -FFTtest/Qwen3-30B-A3B/gdb_sigsegv.gdb -FFTtest/Qwen3-30B-A3B/run_full_ft_test_1gpu_bf16_frozen.sh --gdb -FFTtest/Qwen3-30B-A3B/expert_buf_probe.py -``` - -调试时优先读取原始 `phase4/train.log`、GDB backtrace、逐步 timing 和 probe JSON;自动 -生成的 `summary.md` 只能作为索引,不能覆盖原始证据。 From b6f7a2e68fca714c217f265af60498abdf544df5 Mon Sep 17 00:00:00 2001 From: illu Date: Wed, 15 Jul 2026 09:12:40 +0000 Subject: [PATCH 08/20] [fix](kt-kernel): configure SFT OpenMP threads --- kt-kernel/python/sft/config.py | 85 +++++++++++++++++++ .../test/per_commit/test_sft_omp_threads.py | 67 +++++++++++++++ 2 files changed, 152 insertions(+) create mode 100644 kt-kernel/test/per_commit/test_sft_omp_threads.py diff --git a/kt-kernel/python/sft/config.py b/kt-kernel/python/sft/config.py index 0172ad5dd..a4eca76da 100644 --- a/kt-kernel/python/sft/config.py +++ b/kt-kernel/python/sft/config.py @@ -12,11 +12,18 @@ from __future__ import annotations import dataclasses +import logging import os from dataclasses import dataclass, field +from pathlib import Path from typing import Any, Callable +logger = logging.getLogger(__name__) + +_CPU_TOPOLOGY_ROOT = Path("/sys/devices/system/cpu") + + def _env_int(key: str, default: int | None) -> int | None: value = os.environ.get(key, None) if value is None or value == "": @@ -38,6 +45,83 @@ def _env_bool(key: str, default: bool) -> bool: return value.lower() in ("1", "true", "yes") +def _available_cpu_ids() -> set[int]: + """Return CPUs available to this process, respecting affinity/cpuset limits.""" + try: + return set(os.sched_getaffinity(0)) + except (AttributeError, OSError): + return set(range(os.cpu_count() or 1)) + + +def _read_cpu_topology(cpu_id: int) -> tuple[int, int] | None: + topology = _CPU_TOPOLOGY_ROOT / f"cpu{cpu_id}" / "topology" + try: + package_id = int((topology / "physical_package_id").read_text().strip()) + core_id = int((topology / "core_id").read_text().strip()) + except (OSError, ValueError): + return None + return package_id, core_id + + +def detect_physical_cpu_count() -> int: + """Count physical cores available to the current process. + + Linux exposes a stable ``(physical_package_id, core_id)`` pair for every + logical CPU. Counting those pairs avoids assigning one OpenMP worker to + each SMT sibling. If topology is unavailable, fall back to the number of + affinity-visible logical CPUs. + """ + cpu_ids = _available_cpu_ids() + physical_cores = { + topology + for cpu_id in cpu_ids + if (topology := _read_cpu_topology(cpu_id)) is not None + } + return max(1, len(physical_cores) if physical_cores else len(cpu_ids)) + + +def _set_torch_num_threads(num_threads: int) -> None: + try: + import torch + except ImportError: + return + torch.set_num_threads(num_threads) + + +def configure_omp_threads() -> int: + """Configure OpenMP for KT SFT CPU tensor work. + + ``accelerate launch`` defaults GPU jobs to ``OMP_NUM_THREADS=1`` when the + caller did not choose a value. That makes Full-FT CPU gradient accumulation, + AdamW, and zeroing effectively serial. Treat that value as the launcher + default and select the affinity-visible physical core count instead. + + ``ACCELERATE_KT_OMP_NUM_THREADS`` is the unambiguous KT-specific override, + including when an intentional single-thread run is required. An existing + generic ``OMP_NUM_THREADS`` value greater than one is also preserved. + """ + kt_override = _env_int("ACCELERATE_KT_OMP_NUM_THREADS", None) + current_omp = _env_int("OMP_NUM_THREADS", None) + + if kt_override is not None: + num_threads = kt_override + source = "ACCELERATE_KT_OMP_NUM_THREADS" + elif current_omp is not None and current_omp > 1: + num_threads = current_omp + source = "OMP_NUM_THREADS" + else: + num_threads = detect_physical_cpu_count() + source = "available physical cores" + + if num_threads < 1: + raise ValueError(f"OpenMP thread count must be positive, got {num_threads}") + + os.environ["OMP_NUM_THREADS"] = str(num_threads) + _set_torch_num_threads(num_threads) + logger.info("KT SFT configured OMP_NUM_THREADS=%d from %s", num_threads, source) + return num_threads + + @dataclass class KTConfig: """ @@ -104,6 +188,7 @@ def from_object(cls, obj: Any) -> "KTConfig": return cls(**kwargs) def __post_init__(self): + configure_omp_threads() if self.kt_backend is None: self.kt_backend = os.environ.get("ACCELERATE_KT_BACKEND", "AMXBF16") if self.kt_num_threads is None: diff --git a/kt-kernel/test/per_commit/test_sft_omp_threads.py b/kt-kernel/test/per_commit/test_sft_omp_threads.py new file mode 100644 index 000000000..e180614dc --- /dev/null +++ b/kt-kernel/test/per_commit/test_sft_omp_threads.py @@ -0,0 +1,67 @@ +# SPDX-License-Identifier: Apache-2.0 + +import importlib.util +import os +import sys +from pathlib import Path +from unittest.mock import patch + + +CONFIG_PATH = Path(__file__).resolve().parents[2] / "python" / "sft" / "config.py" +SPEC = importlib.util.spec_from_file_location("kt_sft_config_under_test", CONFIG_PATH) +assert SPEC is not None and SPEC.loader is not None +config = importlib.util.module_from_spec(SPEC) +sys.modules[SPEC.name] = config +SPEC.loader.exec_module(config) + + +def test_detect_physical_cpu_count_deduplicates_smt_siblings(): + topology = { + 0: (0, 0), + 1: (0, 1), + 2: (0, 0), + 3: (0, 1), + 4: (1, 0), + 5: (1, 0), + } + with ( + patch.object(config, "_available_cpu_ids", return_value=set(topology)), + patch.object(config, "_read_cpu_topology", side_effect=topology.get), + ): + assert config.detect_physical_cpu_count() == 3 + + +def test_configure_omp_threads_replaces_accelerate_single_thread_default(): + with ( + patch.dict(os.environ, {"OMP_NUM_THREADS": "1"}, clear=False), + patch.object(config, "detect_physical_cpu_count", return_value=96), + patch.object(config, "_set_torch_num_threads") as set_torch_threads, + ): + os.environ.pop("ACCELERATE_KT_OMP_NUM_THREADS", None) + assert config.configure_omp_threads() == 96 + assert os.environ["OMP_NUM_THREADS"] == "96" + set_torch_threads.assert_called_once_with(96) + + +def test_configure_omp_threads_preserves_explicit_generic_value(): + with ( + patch.dict(os.environ, {"OMP_NUM_THREADS": "48"}, clear=False), + patch.object(config, "_set_torch_num_threads") as set_torch_threads, + ): + os.environ.pop("ACCELERATE_KT_OMP_NUM_THREADS", None) + assert config.configure_omp_threads() == 48 + set_torch_threads.assert_called_once_with(48) + + +def test_configure_omp_threads_supports_explicit_single_thread_override(): + with ( + patch.dict( + os.environ, + {"OMP_NUM_THREADS": "96", "ACCELERATE_KT_OMP_NUM_THREADS": "1"}, + clear=False, + ), + patch.object(config, "_set_torch_num_threads") as set_torch_threads, + ): + assert config.configure_omp_threads() == 1 + assert os.environ["OMP_NUM_THREADS"] == "1" + set_torch_threads.assert_called_once_with(1) From f2098786f02ae4cd1d3f6a2968e15df7dfdf83fc Mon Sep 17 00:00:00 2001 From: illu Date: Thu, 16 Jul 2026 08:12:31 +0000 Subject: [PATCH 09/20] [perf](kt-kernel): optimize AMX Full-FT weight gradients Coarsen base-weight gradient work from individual output tiles to fixed-intermediate strips so each task reuses packed panels across the hidden dimension. Keep aligned thread-local BF16 panels across tasks and retain FP32 AMX accumulator tiles for the full K reduction. Gate and up run separate K passes while sharing the packed input panel. On the matched Qwen3-30B-A3B 1-GPU test, stable Full-FT backward drops from 9.281s to 6.792s (-26.82%), step time drops from 19.252s to 16.140s, and TPS rises from 212.76 to 253.78. The LoRA-only backward control changes by -2.74%. Validated with clang-format, the Release AMX/CUDA extension build, TP1/TP2 reference gradients across boundary token counts, and the 15-step Full-then-LoRA performance run. --- kt-kernel/operators/amx/sft_moe.hpp | 199 ++++++++++++++++------------ 1 file changed, 113 insertions(+), 86 deletions(-) diff --git a/kt-kernel/operators/amx/sft_moe.hpp b/kt-kernel/operators/amx/sft_moe.hpp index aa637cc4c..2e7750f78 100644 --- a/kt-kernel/operators/amx/sft_moe.hpp +++ b/kt-kernel/operators/amx/sft_moe.hpp @@ -14,6 +14,7 @@ #include #include #include +#include #include #include #include @@ -1954,133 +1955,159 @@ class AMX_SFT_MOE_TP : public BaseMOE { const int i_tiles = (I + TILE_M - 1) / TILE_M; const int h_tiles = (H + TILE_N - 1) / TILE_N; - const int tiles_per_projection = i_tiles * h_tiles; - const int tasks_per_expert = tiles_per_projection * 2; // fused gate/up plus down + // Keep enough fixed-i strips for load balancing while amortizing task pickup and panel packing over H. + const int tasks_per_expert = i_tiles * 2; // fixed-i strips for fused gate/up plus down const int total_tasks = activated_expert * tasks_per_expert; auto pool = config_.pool->get_subpool(tp_part_idx); pool->do_work_stealing_job( total_tasks, [](int _) { T::config(); }, - [&, i_tiles, h_tiles, tiles_per_projection, tasks_per_expert](int task_id) { + [&, i_tiles, h_tiles, tasks_per_expert](int task_id) { const int expert_task = task_id / tasks_per_expert; const int local_task = task_id % tasks_per_expert; - const bool do_down = local_task >= tiles_per_projection; - const int tile_id = do_down ? local_task - tiles_per_projection : local_task; + const bool do_down = local_task >= i_tiles; + const int i_tile = local_task % i_tiles; const int expert_idx = cache.m_expert_id_map_cache[expert_task]; const int m = cache.m_local_num_cache[expert_idx]; if (m == 0) return; const size_t pos_start = expert_offsets[expert_task]; - alignas(64) ggml_bf16_t a_tile[TILE_M * TILE_K]; - alignas(64) ggml_bf16_t b_tile[TILE_N * TILE_K]; + const int k_tiles = (m + TILE_K - 1) / TILE_K; + constexpr size_t A_TILE_ELEMENTS = TILE_M * TILE_K; + constexpr size_t B_TILE_ELEMENTS = TILE_N * TILE_K; + constexpr size_t TILE_ALIGNMENT = 64; + constexpr size_t ALIGNMENT_PADDING = TILE_ALIGNMENT / sizeof(ggml_bf16_t); + const size_t packed_a_elements = (size_t)k_tiles * A_TILE_ELEMENTS; + const size_t packed_b_elements = (size_t)k_tiles * B_TILE_ELEMENTS; + + thread_local std::vector packed_a0_storage; + thread_local std::vector packed_a1_storage; + thread_local std::vector packed_b_storage; + auto resize_aligned = [](std::vector& storage, size_t elements) { + const size_t required = elements + ALIGNMENT_PADDING; + if (storage.size() < required) storage.resize(required); + const auto raw = reinterpret_cast(storage.data()); + const auto aligned = (raw + TILE_ALIGNMENT - 1) & ~(std::uintptr_t)(TILE_ALIGNMENT - 1); + return reinterpret_cast(aligned); + }; + + ggml_bf16_t* packed_a0 = resize_aligned(packed_a0_storage, packed_a_elements); + ggml_bf16_t* packed_b = resize_aligned(packed_b_storage, packed_b_elements); alignas(64) float c0[TILE_M * TILE_N]; - alignas(64) float c1[TILE_M * TILE_N]; - if (!do_down) { - const int i_tile = tile_id / h_tiles; - const int h_tile = tile_id % h_tiles; - const int i_start = i_tile * TILE_M; - const int h_start = h_tile * TILE_N; - const int i_count = std::min(TILE_M, I - i_start); - const int h_count = std::min(TILE_N, H - h_start); - const ggml_bf16_t* input = m_local_input_ptr_[expert_idx]; + const int i_start = i_tile * TILE_N; + const int i_count = std::min(TILE_N, I - i_start); - for (int k_start = 0; k_start < m; k_start += TILE_K) { + if (do_down) { + const ggml_bf16_t* grad_output = base_grad_output_bf16_ptr_[expert_idx]; + const ggml_bf16_t* intermediate = cache.intermediate_cache + pos_start * I; + std::memset(packed_b, 0, packed_b_elements * sizeof(ggml_bf16_t)); + for (int kt = 0; kt < k_tiles; kt++) { + const int k_start = kt * TILE_K; const int k_count = std::min(TILE_K, m - k_start); - std::memset(a_tile, 0, sizeof(a_tile)); - std::memset(b_tile, 0, sizeof(b_tile)); - - for (int row = 0; row < i_count; row++) { + ggml_bf16_t* b_tile = packed_b + (size_t)kt * B_TILE_ELEMENTS; + for (int col = 0; col < i_count; col++) { for (int kk = 0; kk < k_count; kk++) { - a_tile[row * TILE_K + kk] = grad_gate_output_[(pos_start + k_start + kk) * I + i_start + row]; - } - } - for (int col = 0; col < h_count; col++) { - for (int kk = 0; kk < k_count; kk++) { - b_tile[col * TILE_K + kk] = input[(size_t)(k_start + kk) * H + h_start + col]; + b_tile[col * TILE_K + kk] = intermediate[(size_t)(k_start + kk) * I + i_start + col]; } } amx::transpose_16x16_32bit(reinterpret_cast<__m512i*>(b_tile)); amx::transpose_16x16_32bit(reinterpret_cast<__m512i*>(b_tile + amx::GemmKernel224BF::TILE_N * TILE_K)); + } - T::load_b(b_tile, TILE_K * sizeof(ggml_bf16_t)); - T::load_a(a_tile, TILE_K * sizeof(ggml_bf16_t)); - if (k_start == 0) { - T::clean_c(); - } else { - T::load_c(c0, TILE_N * sizeof(float)); - } - T::run_tile(); - T::store_c(c0, TILE_N * sizeof(float)); - - std::memset(a_tile, 0, sizeof(a_tile)); - for (int row = 0; row < i_count; row++) { - for (int kk = 0; kk < k_count; kk++) { - a_tile[row * TILE_K + kk] = grad_up_output_[(pos_start + k_start + kk) * I + i_start + row]; + ggml_bf16_t* down_dst = gdp + (size_t)expert_idx * H * F; + for (int h_tile = 0; h_tile < h_tiles; h_tile++) { + const int h_start = h_tile * TILE_M; + const int h_count = std::min(TILE_M, H - h_start); + std::memset(packed_a0, 0, packed_a_elements * sizeof(ggml_bf16_t)); + for (int kt = 0; kt < k_tiles; kt++) { + const int k_start = kt * TILE_K; + const int k_count = std::min(TILE_K, m - k_start); + ggml_bf16_t* a_tile = packed_a0 + (size_t)kt * A_TILE_ELEMENTS; + for (int row = 0; row < h_count; row++) { + for (int kk = 0; kk < k_count; kk++) { + a_tile[row * TILE_K + kk] = grad_output[(size_t)(k_start + kk) * H + h_start + row]; + } } } - T::load_a(a_tile, TILE_K * sizeof(ggml_bf16_t)); - if (k_start == 0) { - T::clean_c(); - } else { - T::load_c(c1, TILE_N * sizeof(float)); + + // Keep the full 32x32 FP32 C tile resident for the complete K reduction. + T::clean_c(); + for (int kt = 0; kt < k_tiles; kt++) { + T::load_b(packed_b + (size_t)kt * B_TILE_ELEMENTS, TILE_K * sizeof(ggml_bf16_t)); + T::load_a(packed_a0 + (size_t)kt * A_TILE_ELEMENTS, TILE_K * sizeof(ggml_bf16_t)); + T::run_tile(); } - T::run_tile(); - T::store_c(c1, TILE_N * sizeof(float)); - } + T::store_c(c0, TILE_N * sizeof(float)); - ggml_bf16_t* gate_dst = ggp + (size_t)expert_idx * F * H; - ggml_bf16_t* up_dst = gup_ptr + (size_t)expert_idx * F * H; - for (int row = 0; row < i_count; row++) { - for (int col = 0; col < h_count; col++) { - gate_dst[(size_t)(i_start + row) * H + h_start + col] = GGML_FP32_TO_BF16(c0[row * TILE_N + col]); - up_dst[(size_t)(i_start + row) * H + h_start + col] = GGML_FP32_TO_BF16(c1[row * TILE_N + col]); + for (int row = 0; row < h_count; row++) { + for (int col = 0; col < i_count; col++) { + down_dst[(size_t)(h_start + row) * F + i_start + col] = GGML_FP32_TO_BF16(c0[row * TILE_N + col]); + } } } return; } - const int h_tile = tile_id / i_tiles; - const int i_tile = tile_id % i_tiles; - const int h_start = h_tile * TILE_M; - const int i_start = i_tile * TILE_N; - const int h_count = std::min(TILE_M, H - h_start); - const int i_count = std::min(TILE_N, I - i_start); - const ggml_bf16_t* grad_output = base_grad_output_bf16_ptr_[expert_idx]; - const ggml_bf16_t* intermediate = cache.intermediate_cache + pos_start * I; - - for (int k_start = 0; k_start < m; k_start += TILE_K) { + ggml_bf16_t* packed_a1 = resize_aligned(packed_a1_storage, packed_a_elements); + std::memset(packed_a0, 0, packed_a_elements * sizeof(ggml_bf16_t)); + std::memset(packed_a1, 0, packed_a_elements * sizeof(ggml_bf16_t)); + for (int kt = 0; kt < k_tiles; kt++) { + const int k_start = kt * TILE_K; const int k_count = std::min(TILE_K, m - k_start); - std::memset(a_tile, 0, sizeof(a_tile)); - std::memset(b_tile, 0, sizeof(b_tile)); - for (int row = 0; row < h_count; row++) { + ggml_bf16_t* gate_a_tile = packed_a0 + (size_t)kt * A_TILE_ELEMENTS; + ggml_bf16_t* up_a_tile = packed_a1 + (size_t)kt * A_TILE_ELEMENTS; + for (int row = 0; row < i_count; row++) { for (int kk = 0; kk < k_count; kk++) { - a_tile[row * TILE_K + kk] = grad_output[(size_t)(k_start + kk) * H + h_start + row]; + gate_a_tile[row * TILE_K + kk] = grad_gate_output_[(pos_start + k_start + kk) * I + i_start + row]; + up_a_tile[row * TILE_K + kk] = grad_up_output_[(pos_start + k_start + kk) * I + i_start + row]; } } - for (int col = 0; col < i_count; col++) { - for (int kk = 0; kk < k_count; kk++) { - b_tile[col * TILE_K + kk] = intermediate[(size_t)(k_start + kk) * I + i_start + col]; + } + + const ggml_bf16_t* input = m_local_input_ptr_[expert_idx]; + ggml_bf16_t* gate_dst = ggp + (size_t)expert_idx * F * H; + ggml_bf16_t* up_dst = gup_ptr + (size_t)expert_idx * F * H; + alignas(64) float c1[TILE_M * TILE_N]; + for (int h_tile = 0; h_tile < h_tiles; h_tile++) { + const int h_start = h_tile * TILE_N; + const int h_count = std::min(TILE_N, H - h_start); + std::memset(packed_b, 0, packed_b_elements * sizeof(ggml_bf16_t)); + for (int kt = 0; kt < k_tiles; kt++) { + const int k_start = kt * TILE_K; + const int k_count = std::min(TILE_K, m - k_start); + ggml_bf16_t* b_tile = packed_b + (size_t)kt * B_TILE_ELEMENTS; + for (int col = 0; col < h_count; col++) { + for (int kk = 0; kk < k_count; kk++) { + b_tile[col * TILE_K + kk] = input[(size_t)(k_start + kk) * H + h_start + col]; + } } + amx::transpose_16x16_32bit(reinterpret_cast<__m512i*>(b_tile)); + amx::transpose_16x16_32bit(reinterpret_cast<__m512i*>(b_tile + amx::GemmKernel224BF::TILE_N * TILE_K)); } - amx::transpose_16x16_32bit(reinterpret_cast<__m512i*>(b_tile)); - amx::transpose_16x16_32bit(reinterpret_cast<__m512i*>(b_tile + amx::GemmKernel224BF::TILE_N * TILE_K)); - T::load_b(b_tile, TILE_K * sizeof(ggml_bf16_t)); - T::load_a(a_tile, TILE_K * sizeof(ggml_bf16_t)); - if (k_start == 0) { - T::clean_c(); - } else { - T::load_c(c0, TILE_N * sizeof(float)); + // Gate and up each consume all four C tiles, so retain C across K in two separate passes. + T::clean_c(); + for (int kt = 0; kt < k_tiles; kt++) { + T::load_b(packed_b + (size_t)kt * B_TILE_ELEMENTS, TILE_K * sizeof(ggml_bf16_t)); + T::load_a(packed_a0 + (size_t)kt * A_TILE_ELEMENTS, TILE_K * sizeof(ggml_bf16_t)); + T::run_tile(); } - T::run_tile(); T::store_c(c0, TILE_N * sizeof(float)); - } - ggml_bf16_t* down_dst = gdp + (size_t)expert_idx * H * F; - for (int row = 0; row < h_count; row++) { - for (int col = 0; col < i_count; col++) { - down_dst[(size_t)(h_start + row) * F + i_start + col] = GGML_FP32_TO_BF16(c0[row * TILE_N + col]); + T::clean_c(); + for (int kt = 0; kt < k_tiles; kt++) { + T::load_b(packed_b + (size_t)kt * B_TILE_ELEMENTS, TILE_K * sizeof(ggml_bf16_t)); + T::load_a(packed_a1 + (size_t)kt * A_TILE_ELEMENTS, TILE_K * sizeof(ggml_bf16_t)); + T::run_tile(); + } + T::store_c(c1, TILE_N * sizeof(float)); + + for (int row = 0; row < i_count; row++) { + for (int col = 0; col < h_count; col++) { + gate_dst[(size_t)(i_start + row) * H + h_start + col] = GGML_FP32_TO_BF16(c0[row * TILE_N + col]); + up_dst[(size_t)(i_start + row) * H + h_start + col] = GGML_FP32_TO_BF16(c1[row * TILE_N + col]); + } } } }, From 9dc9d93b7bf1b00124475cde5ff6cae97741e28d Mon Sep 17 00:00:00 2001 From: yyj Date: Thu, 16 Jul 2026 16:55:18 +0800 Subject: [PATCH 10/20] [feat](kt-kernel): add staged SFT profiling --- kt-kernel/ext_bindings.cpp | 2 + kt-kernel/operators/amx/sft_moe.hpp | 73 ++++++- kt-kernel/operators/moe-sft-tp.hpp | 38 ++++ kt-kernel/operators/sft_profile.hpp | 200 ++++++++++++++++++ kt-kernel/python/sft/__init__.py | 4 + kt-kernel/python/sft/profiler.py | 143 +++++++++++++ .../test/per_commit/test_sft_profiler.py | 68 ++++++ 7 files changed, 527 insertions(+), 1 deletion(-) create mode 100644 kt-kernel/operators/sft_profile.hpp create mode 100644 kt-kernel/python/sft/profiler.py create mode 100644 kt-kernel/test/per_commit/test_sft_profiler.py diff --git a/kt-kernel/ext_bindings.cpp b/kt-kernel/ext_bindings.cpp index 95557973e..f3408f935 100644 --- a/kt-kernel/ext_bindings.cpp +++ b/kt-kernel/ext_bindings.cpp @@ -411,6 +411,8 @@ void bind_moe_sft_module(py::module_& moe_module, const char* name) { }) .def("submit_backward_repack", &MoeClass::submit_backward_repack) .def("wait_backward_repack", &MoeClass::wait_backward_repack) + .def("get_profile_stats", &MoeClass::get_profile_stats, py::arg("reset") = false) + .def("reset_profile_stats", &MoeClass::reset_profile_stats) // Update base weight BF16 pointers for reload_base_weights (full mode training) // After calling this, call load_weights_task() to re-quantize BF16->AMX .def("set_base_weight_pointers", diff --git a/kt-kernel/operators/amx/sft_moe.hpp b/kt-kernel/operators/amx/sft_moe.hpp index 2e7750f78..d0deb24b5 100644 --- a/kt-kernel/operators/amx/sft_moe.hpp +++ b/kt-kernel/operators/amx/sft_moe.hpp @@ -27,6 +27,7 @@ #include #include "../../cpu_backend/worker_pool.h" +#include "../sft_profile.hpp" #include "ggml.h" #include "la/amx_kernels.hpp" #include "la/avx_kernels.hpp" @@ -435,6 +436,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { // SFT configuration MOESFTConfig sft_config_; + SFTProfiler profiler_; // LoRA configuration (from MOESFTConfig) int lora_rank_; @@ -703,6 +705,13 @@ class AMX_SFT_MOE_TP : public BaseMOE { free_transposed_lora_weights(); } + void append_profile_stats(std::map& out, const std::string& prefix, + bool reset_after = false) { + profiler_.append(out, prefix, reset_after); + } + + void reset_profile_stats() { profiler_.reset(); } + /** * @brief Allocate forward-phase buffers. * Called at the start of forward_sft. @@ -934,7 +943,8 @@ class AMX_SFT_MOE_TP : public BaseMOE { */ void forward_sft(int qlen, int k, const int64_t* expert_ids, const float* weights, const void* input, void* output, bool save_for_backward) { - uint64_t _fwd_start_cycles = __rdtsc(); + SFTProfileScope total_scope(profiler_, SFTProfileStage::FwdTotal); + auto stage_start = profiler_.start(); SFT_POOL_LOG("fwd_enter", config_.layer_idx, tp_part_idx, qlen, cache_stack_top_, forward_pool_bytes_, cache_pool_bytes_, backward_pool_bytes_, 0, "save_bwd=%d", (int)save_for_backward); @@ -966,8 +976,10 @@ class AMX_SFT_MOE_TP : public BaseMOE { transpose_lora_b_weights(); lora_b_transposed_ = true; } + profiler_.record(SFTProfileStage::FwdSetup, stage_start); // Step 1: Expert routing (reuse base class logic) + stage_start = profiler_.start(); int activated_expert = 0; std::fill(m_local_num_.begin(), m_local_num_.end(), 0); for (int i = 0; i < qlen; i++) { @@ -985,8 +997,12 @@ class AMX_SFT_MOE_TP : public BaseMOE { activated_expert++; } } + profiler_.record(SFTProfileStage::FwdRoute, stage_start); + profiler_.record_workload(static_cast(qlen), static_cast(qlen) * k, + static_cast(activated_expert)); // Step 2: Buffer pool allocation (reuse base class logic) + stage_start = profiler_.start(); size_t offset = 0; void* gate_up_ba_pool_ptr = Base::gate_up_ba_pool_; void* gate_bc_pool_ptr = Base::gate_bc_pool_; @@ -1087,6 +1103,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { printf("[OVERFLOW DEBUG L%d] Total tokens processed: %zu (offset after loop)\n", config_.layer_idx, offset); } } + profiler_.record(SFTProfileStage::FwdBufferSetup, stage_start); // Step 3: Copy input to expert buffers auto direct_or_pool = [&](int count, auto&& fn) { @@ -1099,6 +1116,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { } }; + stage_start = profiler_.start(); direct_or_pool(qlen, [&](int i) { for (int j = 0; j < k; j++) { if (expert_ids[i * k + j] < config_.num_gpu_experts || expert_ids[i * k + j] >= config_.expert_num) { @@ -1108,6 +1126,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { (ggml_bf16_t*)input + i * config_.hidden_size, sizeof(ggml_bf16_t) * config_.hidden_size); } }); + profiler_.record(SFTProfileStage::FwdInputScatter, stage_start); // NaN Check: Step 3 - Packed input if (is_nan_check_enabled()) { @@ -1124,12 +1143,15 @@ class AMX_SFT_MOE_TP : public BaseMOE { } // Step 4: Quantize input + stage_start = profiler_.start(); direct_or_pool(activated_expert, [this](int task_id) { int expert_idx = m_expert_id_map_[task_id]; gate_up_ba_[expert_idx]->from_mat(m_local_num_[expert_idx], m_local_input_ptr_[expert_idx], 0, 1); }); + profiler_.record(SFTProfileStage::FwdInputPack, stage_start); // Step 5: Gate + Up GEMM (base projection) + stage_start = profiler_.start(); int nth = T::recommended_nth(config_.intermediate_size); pool->do_work_stealing_job( nth * activated_expert * 2, [](int _) { T::config(); }, @@ -1146,6 +1168,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { } }, nullptr); + profiler_.record(SFTProfileStage::FwdGateUpBase, stage_start); // NaN Check: Step 5 - Gate/Up GEMM output (before LoRA) if (is_nan_check_enabled()) { @@ -1166,9 +1189,11 @@ class AMX_SFT_MOE_TP : public BaseMOE { } // Step 5.5: Gate + Up LoRA (AVX512 BF16 - no BufferB conversion needed) + stage_start = profiler_.start(); if (!SkipLoRA) { compute_lora_gate_up(qlen, activated_expert); } + profiler_.record(SFTProfileStage::FwdGateUpLora, stage_start); // NaN Check: Step 5.5 - Gate/Up output (after LoRA) if (is_nan_check_enabled()) { @@ -1189,6 +1214,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { } // Save gate/up outputs before activation (for backward) + stage_start = profiler_.start(); if (save_for_backward) { // If a cache entry already exists (checkpoint recompute scenario), // overwrite it instead of pushing a new one. This keeps the cache @@ -1233,9 +1259,12 @@ class AMX_SFT_MOE_TP : public BaseMOE { check_cache_bf16("up_output_cache", cache.up_output_cache, total_tokens * config_.intermediate_size); } } + profiler_.record(SFTProfileStage::FwdCacheGateUp, stage_start); // Step 6: Activation (silu(gate) * up) + stage_start = profiler_.start(); { Base::apply_activation(activated_expert, nth, qlen); } + profiler_.record(SFTProfileStage::FwdActivation, stage_start); // NaN Check: Step 6 - Activation output (silu(gate) * up) if (is_nan_check_enabled()) { @@ -1252,6 +1281,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { } // Save intermediate AFTER activation for backward_down (Bug #17c fix) + stage_start = profiler_.start(); if (save_for_backward) { ForwardCache& cache = cache_stack_[cache_stack_top_ - 1]; // Get the cache we just pushed save_intermediate_to_cache(cache, activated_expert); @@ -1288,8 +1318,10 @@ class AMX_SFT_MOE_TP : public BaseMOE { } } } + profiler_.record(SFTProfileStage::FwdCacheIntermediate, stage_start); // Step 7: Quantize intermediate for down projection + stage_start = profiler_.start(); pool->do_work_stealing_job( activated_expert, nullptr, [this](int task_id) { @@ -1297,8 +1329,10 @@ class AMX_SFT_MOE_TP : public BaseMOE { down_ba_[expert_idx]->from_mat(m_local_num_[expert_idx], m_local_gate_output_ptr_[expert_idx], 0, 1); }, nullptr); + profiler_.record(SFTProfileStage::FwdDownPack, stage_start); // Step 8: Down GEMM + stage_start = profiler_.start(); nth = T::recommended_nth(config_.hidden_size); pool->do_work_stealing_job( nth * activated_expert, [](int _) { T::config(); }, @@ -1309,6 +1343,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { down_bc_[expert_idx]->to_mat(m_local_num_[expert_idx], m_local_down_output_ptr_[expert_idx], ith, nth); }, nullptr); + profiler_.record(SFTProfileStage::FwdDownBase, stage_start); // NaN Check: Step 8 - Down GEMM output (before LoRA) if (is_nan_check_enabled()) { @@ -1325,10 +1360,12 @@ class AMX_SFT_MOE_TP : public BaseMOE { } // Step 8.5: Down LoRA (AVX512 BF16 - no BufferB conversion needed) + stage_start = profiler_.start(); if (down_lora_a_ != nullptr && down_lora_b_ != nullptr) { ForwardCache* cache_ptr = save_for_backward ? &cache_stack_[cache_stack_top_ - 1] : nullptr; compute_lora_down(qlen, activated_expert, cache_ptr); } + profiler_.record(SFTProfileStage::FwdDownLora, stage_start); // NaN Check: Step 8.5 - Down output (after LoRA) if (is_nan_check_enabled()) { @@ -1345,12 +1382,15 @@ class AMX_SFT_MOE_TP : public BaseMOE { } // Save down_output for grad_weights computation + stage_start = profiler_.start(); if (save_for_backward) { ForwardCache& cache = cache_stack_[cache_stack_top_ - 1]; // Get the cache we just pushed save_down_output_to_cache(cache, activated_expert); } + profiler_.record(SFTProfileStage::FwdCacheDown, stage_start); // Step 9: Weighted merge + stage_start = profiler_.start(); pool->do_work_stealing_job( qlen, nullptr, [this, output, k, expert_ids, weights](int i) { @@ -1375,6 +1415,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { } }, nullptr); + profiler_.record(SFTProfileStage::FwdWeightedMerge, stage_start); // NaN Check: Step 9 - Final output (after weighted merge) if (is_nan_check_enabled()) { @@ -1414,6 +1455,8 @@ class AMX_SFT_MOE_TP : public BaseMOE { void* grad_weights, int full_intermediate_size = 0, float* fp32_grad_down_lora_b = nullptr, float* fp32_grad_gate_lora_a = nullptr, float* fp32_grad_up_lora_a = nullptr, void* grad_gate_proj = nullptr, void* grad_up_proj = nullptr, void* grad_down_proj = nullptr) { + SFTProfileScope total_scope(profiler_, SFTProfileStage::BwdTotal); + auto stage_start = profiler_.start(); // If full_intermediate_size not provided, use local (non-TP mode) if (full_intermediate_size == 0) full_intermediate_size = config_.intermediate_size; SFT_POOL_LOG("bwd_enter", config_.layer_idx, tp_part_idx, 0, cache_stack_top_, forward_pool_bytes_, @@ -1539,10 +1582,12 @@ class AMX_SFT_MOE_TP : public BaseMOE { m_local_down_output_ptr_[i] = m_local_down_output_ + offset * config_.hidden_size; offset += m_local_num_[i]; } + profiler_.record(SFTProfileStage::BwdSetup, stage_start); // Restore input data from cache into m_local_input_ (shared_mem_buffer may have been // overwritten by subsequent layers' forward passes). This is needed for gate/up LoRA // gradient computation which reads from m_local_input_ptr_. + stage_start = profiler_.start(); auto pool_local = config_.pool->get_subpool(tp_part_idx); auto restore_input = [&](int i) { for (int j = 0; j < k; j++) { @@ -1564,14 +1609,17 @@ class AMX_SFT_MOE_TP : public BaseMOE { } else { pool_local->do_work_stealing_job(qlen, nullptr, restore_input, nullptr); } + profiler_.record(SFTProfileStage::BwdCacheRestore, stage_start); // Step 1: Down projection backward + stage_start = profiler_.start(); if constexpr (supports_standard_mat_mul_v) { backward_down_amx(cache, grad_output, grad_down_lora_a, grad_down_lora_b, full_intermediate_size, fp32_grad_down_lora_b); } else { // backward_down(cache, grad_output, grad_down_lora_a, grad_down_lora_b); } + profiler_.record(SFTProfileStage::BwdDownTotal, stage_start); // // Compute total tokens for debug // size_t total_tokens = 0; @@ -1647,7 +1695,9 @@ class AMX_SFT_MOE_TP : public BaseMOE { // } // } + stage_start = profiler_.start(); backward_activation(cache); + profiler_.record(SFTProfileStage::BwdActivation, stage_start); // NaN Check: Step 2 - After backward_activation if (is_nan_check_enabled()) { @@ -1687,12 +1737,14 @@ class AMX_SFT_MOE_TP : public BaseMOE { // } // } + stage_start = profiler_.start(); if constexpr (supports_standard_mat_mul_v) { backward_gate_up_amx(cache, grad_input, grad_gate_lora_a, grad_gate_lora_b, grad_up_lora_a, grad_up_lora_b, full_intermediate_size, fp32_grad_gate_lora_a, fp32_grad_up_lora_a); } else { // backward_gate_up(cache, grad_input, grad_gate_lora_a, grad_gate_lora_b, grad_up_lora_a, grad_up_lora_b); } + profiler_.record(SFTProfileStage::BwdGateUpTotal, stage_start); // NaN Check: Step 3 - After backward_gate_up if (is_nan_check_enabled()) { @@ -1729,6 +1781,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { // Step 4: Compute grad_weights (gradient for routing weights) // grad_weights[token_idx, expert_pos] = dot(grad_output[token_idx], down_output[token, expert]) + stage_start = profiler_.start(); if (grad_weights != nullptr) { auto pool = config_.pool->get_subpool(tp_part_idx); float* grad_w = (float*)grad_weights; @@ -1780,6 +1833,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { pool->do_work_stealing_job(qlen, nullptr, compute_grad_weight, nullptr); } } + profiler_.record(SFTProfileStage::BwdRouterGrad, stage_start); // NaN Check: Step 4 - After grad_weights computation if (is_nan_check_enabled() && grad_weights != nullptr) { @@ -1900,9 +1954,11 @@ class AMX_SFT_MOE_TP : public BaseMOE { // ===================================================================== // Step 5: Base weight gradient accumulation (full weight grad mode) // ===================================================================== + stage_start = profiler_.start(); if (sft_config_.full_weight_grad && grad_gate_proj && grad_up_proj && grad_down_proj) { backward_base_weight_grad(cache, full_intermediate_size, grad_gate_proj, grad_up_proj, grad_down_proj); } + profiler_.record(SFTProfileStage::BwdBaseWeightGrad, stage_start); // \u2605 Cache pool is NOT freed here \u2014 kept for reuse across steps. // alloc_or_resize_cache_pool() is grow-only, so same-seqlen steps @@ -2606,6 +2662,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { * and sets the owner layer on the shared pool. */ void prepare_backward_bb_for_async() { + SFTProfileScope profile_scope(profiler_, SFTProfileStage::BackwardRepack); if constexpr (!supports_standard_mat_mul_v) return; if (backward_bb_pool_bytes_ == 0) return; @@ -4452,6 +4509,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { void backward_down_amx(const ForwardCache& cache, const void* grad_output, void* grad_down_lora_a, void* grad_down_lora_b, int full_intermediate_size = 0, float* fp32_grad_down_lora_b = nullptr) { + auto stage_start = profiler_.start(); if (full_intermediate_size == 0) full_intermediate_size = config_.intermediate_size; auto pool = config_.pool->get_subpool(tp_part_idx); int activated_expert = cache.activated_expert_cache; @@ -4514,6 +4572,8 @@ class AMX_SFT_MOE_TP : public BaseMOE { } // NOTE: no full-buffer memset here; grad_intermediate_ is overwritten by to_mat() for active tokens. + profiler_.record(SFTProfileStage::BwdDownSetup, stage_start); + stage_start = profiler_.start(); // ===================================================== // Step 1: Zero per-expert grad_output buffers @@ -4580,6 +4640,8 @@ class AMX_SFT_MOE_TP : public BaseMOE { static_cast(num_tokens) * config_.hidden_size * sizeof(ggml_bf16_t)); }); } + profiler_.record(SFTProfileStage::BwdDownScatter, stage_start); + stage_start = profiler_.start(); // ===================================================== // Step 3: Quantize scattered grad_output to BufferA @@ -4635,6 +4697,8 @@ class AMX_SFT_MOE_TP : public BaseMOE { bc->to_mat(m, grad_intermediate_ + expert_offsets[task_idx], ith, nth); }, nullptr); + profiler_.record(SFTProfileStage::BwdDownBaseDx, stage_start); + stage_start = profiler_.start(); // ===================================================== // Step 3.5: Add LoRA contribution to grad_intermediate (AVX512) @@ -5076,6 +5140,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { } } } + profiler_.record(SFTProfileStage::BwdDownLora, stage_start); } void backward_activation(const ForwardCache& cache) { @@ -5212,6 +5277,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { void backward_gate_up_amx(const ForwardCache& cache, void* grad_input, void* grad_gate_lora_a, void* grad_gate_lora_b, void* grad_up_lora_a, void* grad_up_lora_b, int full_intermediate_size = 0, float* fp32_grad_gate_lora_a = nullptr, float* fp32_grad_up_lora_a = nullptr) { + auto stage_start = profiler_.start(); if (full_intermediate_size == 0) full_intermediate_size = config_.intermediate_size; auto pool = config_.pool->get_subpool(tp_part_idx); int activated_expert = cache.activated_expert_cache; @@ -5393,8 +5459,11 @@ class AMX_SFT_MOE_TP : public BaseMOE { scatter_to_grad_input(1.0f); }; + profiler_.record(SFTProfileStage::BwdGateUpSetup, stage_start); + stage_start = profiler_.start(); base_pass(false); // gate base_pass(true); // up + profiler_.record(SFTProfileStage::BwdGateUpBaseDx, stage_start); // // DEBUG: Check m_local_input_ptr_ AFTER base_pass (before LoRA) // { @@ -5423,6 +5492,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { if (SkipLoRA || lora_rank_ <= 0 || gate_lora_a_ == nullptr || gate_lora_b_ == nullptr) { return; } + stage_start = profiler_.start(); const bool use_fp32_lora_a = (fp32_grad_gate_lora_a != nullptr); @@ -5790,6 +5860,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { lora_pass_remainder(false); // gate: gb_gradin_fused, scatter, gradA lora_pass_remainder(true); // up: gb_gradin_fused, scatter, gradA + profiler_.record(SFTProfileStage::BwdGateUpLora, stage_start); } }; diff --git a/kt-kernel/operators/moe-sft-tp.hpp b/kt-kernel/operators/moe-sft-tp.hpp index 918aaf74d..79ecadf49 100644 --- a/kt-kernel/operators/moe-sft-tp.hpp +++ b/kt-kernel/operators/moe-sft-tp.hpp @@ -26,6 +26,7 @@ #include "amx/la/amx.hpp" #include "moe-tp.hpp" +#include "sft_profile.hpp" struct TPBf16Stats { double abs_mean = 0.0; @@ -201,6 +202,7 @@ class TP_MOE_SFT : public TP_MOE { // Async backward repack state (Phase 2: overlap repack with GPU attention backward) std::thread repack_thread_; std::atomic repack_in_flight_{false}; + SFTProfiler profiler_; // Per-instance references to shared per-TP backward temporary pools. std::vector backward_temp_pools_; @@ -250,6 +252,22 @@ class TP_MOE_SFT : public TP_MOE { } } + std::map get_profile_stats(bool reset_after = false) { + std::map stats; + stats["layer_idx"] = static_cast(config.layer_idx); + stats["tp_count"] = static_cast(tp_count); + profiler_.append(stats, "wrapper.", reset_after); + for (int i = 0; i < tp_count; ++i) { + tps[i]->append_profile_stats(stats, "tp." + std::to_string(i) + ".", reset_after); + } + return stats; + } + + void reset_profile_stats() { + profiler_.reset(); + for (int i = 0; i < tp_count; ++i) tps[i]->reset_profile_stats(); + } + /** * @brief Load weights on all NUMA nodes with TP partitioning. * @@ -259,6 +277,7 @@ class TP_MOE_SFT : public TP_MOE { * resulting in 2x the expected output after merge. */ void load_weights() override { + SFTProfileScope profile_scope(profiler_, SFTProfileStage::BaseWeightReload); auto pool = config.pool; const uint64_t* physical_to_logical_map = (const uint64_t*)config.physical_to_logical_map; @@ -495,6 +514,7 @@ class TP_MOE_SFT : public TP_MOE { void forward_sft(int* qlen_ptr, int k, const int64_t* expert_ids, const float* weights, const void* input, void* output, bool save_for_backward) { + SFTProfileScope total_scope(profiler_, SFTProfileStage::TpFwdTotal); if (weights_loaded == false) [[unlikely]] { throw std::runtime_error("Weights not loaded"); } @@ -508,10 +528,12 @@ class TP_MOE_SFT : public TP_MOE { } // Run forward on each NUMA node + auto stage_start = profiler_.start(); pool->dispense_backend()->do_numa_job([this, qlen, k, expert_ids, input, weights, save_for_backward](int numa_id) { tps[numa_id]->forward_sft(qlen, k, expert_ids, weights, input, this->local_output_numa[numa_id], save_for_backward); }); + profiler_.record(SFTProfileStage::TpFwdNumaCompute, stage_start); // // Collect per-thread timing from all NUMA subpools // for (int i = 0; i < tp_count; i++) { @@ -520,7 +542,9 @@ class TP_MOE_SFT : public TP_MOE { // // Print per-thread forward timing // Merge results from all NUMA nodes + stage_start = profiler_.start(); this->merge_results(qlen, output); + profiler_.record(SFTProfileStage::TpFwdMerge, stage_start); pool->dispense_backend()->do_numa_job([&](int numa_id) {}); } @@ -556,6 +580,8 @@ class TP_MOE_SFT : public TP_MOE { void* grad_up_lora_a, void* grad_up_lora_b, void* grad_down_lora_a, void* grad_down_lora_b, void* grad_weights, void* grad_gate_proj = nullptr, void* grad_up_proj = nullptr, void* grad_down_proj = nullptr) { + SFTProfileScope total_scope(profiler_, SFTProfileStage::TpBwdTotal); + auto stage_start = profiler_.start(); auto pool = config.pool; // Get full intermediate_size (before TP partitioning) @@ -579,6 +605,8 @@ class TP_MOE_SFT : public TP_MOE { if (active_count > 0) { std::memcpy(active_expert_map.data(), tps[0]->get_cache_expert_id_map(), active_count * sizeof(int)); } + profiler_.record_workload(static_cast(qlen), static_cast(qlen) * k, + static_cast(active_count)); // ===================================================================== // Allocate per-TP temporary buffers. @@ -676,6 +704,7 @@ class TP_MOE_SFT : public TP_MOE { std::memset(seg.ptr, 0, seg.len); }, nullptr); + profiler_.record(SFTProfileStage::TpBwdBufferClear, stage_start); // Compute TP-slice pointers for copy-type direct writes // Each TP writes to its own I-slice of the final output tensor @@ -720,6 +749,7 @@ class TP_MOE_SFT : public TP_MOE { } // Run backward on each NUMA node + stage_start = profiler_.start(); pool->dispense_backend()->do_numa_job([&](int numa_id) { tps[numa_id]->backward(grad_output, part_grad_input_[numa_id], // reduce-type: BF16 pointer unused (FP32 sparse used instead) @@ -733,6 +763,7 @@ class TP_MOE_SFT : public TP_MOE { tp_fp32_gate_a[numa_id], tp_fp32_up_a[numa_id], tp_grad_gate_proj[numa_id], tp_grad_up_proj[numa_id], tp_grad_down_proj[numa_id]); }); + profiler_.record(SFTProfileStage::TpBwdNumaCompute, stage_start); // // Collect per-thread timing from all NUMA subpools // for (int i = 0; i < tp_count; i++) { @@ -758,6 +789,7 @@ class TP_MOE_SFT : public TP_MOE { // } // Bug #22 fix: Merge grad_input from all NUMA nodes (sum them together) + stage_start = profiler_.start(); { auto* out = (ggml_bf16_t*)grad_input; pool->do_work_stealing_job( @@ -804,9 +836,11 @@ class TP_MOE_SFT : public TP_MOE { }, nullptr); } + profiler_.record(SFTProfileStage::TpBwdGradInputMerge, stage_start); // Merge reduce-type LoRA gradients: sparse FP32 sum across TPs → BF16 final output // Copy-type grads (gate/up_lora_b, down_lora_a) were written directly — no merge needed. + stage_start = profiler_.start(); if constexpr (!kSkipLoRA) { // Sparse merge for gate_lora_a, up_lora_a: [active_count, r, H] FP32 → [E, r, H] BF16 { @@ -880,9 +914,11 @@ class TP_MOE_SFT : public TP_MOE { nullptr); } } // if constexpr (!kSkipLoRA) + profiler_.record(SFTProfileStage::TpBwdLoraMerge, stage_start); // Merge grad_weights from all NUMA nodes (sum them together) // Each NUMA computes partial grad_weights based on its down_output partition + stage_start = profiler_.start(); if (grad_weights != nullptr) { float* out_grad_weights = (float*)grad_weights; const size_t total = (size_t)qlen * (size_t)k; @@ -918,6 +954,7 @@ class TP_MOE_SFT : public TP_MOE { }, nullptr); } + profiler_.record(SFTProfileStage::TpBwdRouterGradMerge, stage_start); pool->dispense_backend()->do_numa_job([&](int numa_id) {}); } @@ -1107,6 +1144,7 @@ class TP_MOE_SFT : public TP_MOE { repack_in_flight_.store(true, std::memory_order_release); repack_thread_ = std::thread([this]() { + SFTProfileScope profile_scope(profiler_, SFTProfileStage::BackwardRepack); config.pool->dispense_backend()->do_numa_job( [this](int numa_id) { tps[numa_id]->prepare_backward_bb_for_async(); }); repack_in_flight_.store(false, std::memory_order_release); diff --git a/kt-kernel/operators/sft_profile.hpp b/kt-kernel/operators/sft_profile.hpp new file mode 100644 index 000000000..17dc77198 --- /dev/null +++ b/kt-kernel/operators/sft_profile.hpp @@ -0,0 +1,200 @@ +// Lightweight staged profiling for KT SFT MoE operators. +// SPDX-License-Identifier: Apache-2.0 + +#ifndef CPUINFER_OPERATOR_SFT_PROFILE_HPP +#define CPUINFER_OPERATOR_SFT_PROFILE_HPP + +#include +#include +#include +#include +#include +#include +#include +#include + +enum class SFTProfileStage : uint8_t { + // NUMA-local forward stages. + FwdTotal, + FwdSetup, + FwdRoute, + FwdBufferSetup, + FwdInputScatter, + FwdInputPack, + FwdGateUpBase, + FwdGateUpLora, + FwdCacheGateUp, + FwdActivation, + FwdCacheIntermediate, + FwdDownPack, + FwdDownBase, + FwdDownLora, + FwdCacheDown, + FwdWeightedMerge, + + // NUMA-local backward stages. + BwdTotal, + BwdSetup, + BwdCacheRestore, + BwdDownTotal, + BwdDownSetup, + BwdDownScatter, + BwdDownBaseDx, + BwdDownLora, + BwdActivation, + BwdGateUpTotal, + BwdGateUpSetup, + BwdGateUpBaseDx, + BwdGateUpLora, + BwdRouterGrad, + BwdBaseWeightGrad, + + // TP wrapper and weight-layout stages. + TpFwdTotal, + TpFwdNumaCompute, + TpFwdMerge, + TpBwdTotal, + TpBwdBufferClear, + TpBwdNumaCompute, + TpBwdGradInputMerge, + TpBwdLoraMerge, + TpBwdRouterGradMerge, + BackwardRepack, + BaseWeightReload, + + Count, +}; + +inline constexpr std::array(SFTProfileStage::Count)> kSFTProfileStageNames = { + "forward.total", + "forward.setup", + "forward.route", + "forward.buffer_setup", + "forward.input_scatter", + "forward.input_pack", + "forward.gate_up_base", + "forward.gate_up_lora", + "forward.cache_gate_up", + "forward.activation", + "forward.cache_intermediate", + "forward.down_pack", + "forward.down_base", + "forward.down_lora", + "forward.cache_down", + "forward.weighted_merge", + "backward.total", + "backward.setup", + "backward.cache_restore", + "backward.down.total", + "backward.down.setup", + "backward.down.scatter", + "backward.down.base_dx", + "backward.down.lora", + "backward.activation", + "backward.gate_up.total", + "backward.gate_up.setup", + "backward.gate_up.base_dx", + "backward.gate_up.lora", + "backward.router_grad", + "backward.base_weight_grad", + "tp.forward.total", + "tp.forward.numa_compute", + "tp.forward.merge", + "tp.backward.total", + "tp.backward.buffer_clear", + "tp.backward.numa_compute", + "tp.backward.grad_input_merge", + "tp.backward.lora_merge", + "tp.backward.router_grad_merge", + "weights.backward_repack", + "weights.base_reload", +}; + +inline bool sft_profile_enabled_from_env() { + const char* value = std::getenv("KT_SFT_PROFILE"); + return value != nullptr && value[0] != '\0' && std::strcmp(value, "0") != 0 && + std::strcmp(value, "false") != 0 && std::strcmp(value, "False") != 0; +} + +class SFTProfiler { + public: + using Clock = std::chrono::steady_clock; + using TimePoint = Clock::time_point; + + explicit SFTProfiler(bool enabled = sft_profile_enabled_from_env()) : enabled_(enabled) { + reset(); + } + + bool enabled() const { return enabled_; } + + TimePoint start() const { return enabled_ ? Clock::now() : TimePoint{}; } + + void record(SFTProfileStage stage, TimePoint start) { + if (!enabled_) return; + const auto elapsed = std::chrono::duration_cast(Clock::now() - start).count(); + const size_t idx = static_cast(stage); + total_ns_[idx].fetch_add(static_cast(elapsed), std::memory_order_relaxed); + calls_[idx].fetch_add(1, std::memory_order_relaxed); + } + + void record_workload(uint64_t tokens, uint64_t routed_rows, uint64_t active_experts) { + if (!enabled_) return; + tokens_.fetch_add(tokens, std::memory_order_relaxed); + routed_rows_.fetch_add(routed_rows, std::memory_order_relaxed); + active_experts_.fetch_add(active_experts, std::memory_order_relaxed); + workloads_.fetch_add(1, std::memory_order_relaxed); + } + + void append(std::map& out, const std::string& prefix, bool reset_after = false) { + out[prefix + "enabled"] = enabled_ ? 1.0 : 0.0; + out[prefix + "workloads"] = static_cast(load_or_exchange(workloads_, reset_after)); + out[prefix + "tokens"] = static_cast(load_or_exchange(tokens_, reset_after)); + out[prefix + "routed_rows"] = static_cast(load_or_exchange(routed_rows_, reset_after)); + out[prefix + "active_experts"] = static_cast(load_or_exchange(active_experts_, reset_after)); + for (size_t i = 0; i < static_cast(SFTProfileStage::Count); ++i) { + const std::string stage_prefix = prefix + kSFTProfileStageNames[i] + "."; + out[stage_prefix + "total_ns"] = static_cast(load_or_exchange(total_ns_[i], reset_after)); + out[stage_prefix + "calls"] = static_cast(load_or_exchange(calls_[i], reset_after)); + } + } + + void reset() { + for (auto& value : total_ns_) value.store(0, std::memory_order_relaxed); + for (auto& value : calls_) value.store(0, std::memory_order_relaxed); + workloads_.store(0, std::memory_order_relaxed); + tokens_.store(0, std::memory_order_relaxed); + routed_rows_.store(0, std::memory_order_relaxed); + active_experts_.store(0, std::memory_order_relaxed); + } + + private: + static uint64_t load_or_exchange(std::atomic& value, bool reset_after) { + return reset_after ? value.exchange(0, std::memory_order_relaxed) : value.load(std::memory_order_relaxed); + } + + bool enabled_; + std::array, static_cast(SFTProfileStage::Count)> total_ns_{}; + std::array, static_cast(SFTProfileStage::Count)> calls_{}; + std::atomic workloads_{0}; + std::atomic tokens_{0}; + std::atomic routed_rows_{0}; + std::atomic active_experts_{0}; +}; + +class SFTProfileScope { + public: + SFTProfileScope(SFTProfiler& profiler, SFTProfileStage stage) + : profiler_(profiler), stage_(stage), start_(profiler.start()) {} + + ~SFTProfileScope() { profiler_.record(stage_, start_); } + + SFTProfileScope(const SFTProfileScope&) = delete; + SFTProfileScope& operator=(const SFTProfileScope&) = delete; + + private: + SFTProfiler& profiler_; + SFTProfileStage stage_; + SFTProfiler::TimePoint start_; +}; + +#endif // CPUINFER_OPERATOR_SFT_PROFILE_HPP diff --git a/kt-kernel/python/sft/__init__.py b/kt-kernel/python/sft/__init__.py index 88b266e08..b17ed6873 100644 --- a/kt-kernel/python/sft/__init__.py +++ b/kt-kernel/python/sft/__init__.py @@ -52,6 +52,7 @@ get_kt_loading_kwargs, load_kt_model, ) +from .profiler import collect_kt_sft_profile, format_kt_sft_profile, reset_kt_sft_profile __all__ = [ "KTConfig", @@ -89,4 +90,7 @@ "build_kt_device_map_simplified", "get_kt_loading_kwargs", "load_kt_model", + "collect_kt_sft_profile", + "format_kt_sft_profile", + "reset_kt_sft_profile", ] diff --git a/kt-kernel/python/sft/profiler.py b/kt-kernel/python/sft/profiler.py new file mode 100644 index 000000000..1cab47d26 --- /dev/null +++ b/kt-kernel/python/sft/profiler.py @@ -0,0 +1,143 @@ +# Staged profiling helpers for KT SFT MoE. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from collections import defaultdict +from typing import Any + + +def _find_kt_wrappers(model: Any): + wrappers = getattr(model, "_kt_wrappers", None) + if wrappers is None: + base_model = model + for attr in ("base_model", "model"): + if hasattr(base_model, attr): + base_model = getattr(base_model, attr) + wrappers = getattr(base_model, "_kt_wrappers", None) + if wrappers: + break + return wrappers + + +def _profile_moe(layer_wrapper: Any): + backend = getattr(layer_wrapper, "wrapper", layer_wrapper) + return getattr(backend, "moe", None) + + +def collect_kt_sft_profile(model: Any, reset: bool = False) -> dict[str, Any]: + """Collect JSON-compatible staged timing snapshots from all KT MoE layers.""" + layers: dict[int, dict[str, float]] = {} + enabled = False + for layer_wrapper in _find_kt_wrappers(model) or []: + moe = _profile_moe(layer_wrapper) + if moe is None or not hasattr(moe, "get_profile_stats"): + continue + raw = {str(key): float(value) for key, value in moe.get_profile_stats(reset).items()} + layer_idx = int(raw.get("layer_idx", getattr(layer_wrapper, "layer_idx", len(layers)))) + layers[layer_idx] = raw + enabled = enabled or bool(raw.get("wrapper.enabled", 0.0)) + return {"enabled": enabled, "layers": layers} + + +def reset_kt_sft_profile(model: Any) -> None: + """Reset staged timing counters on all KT MoE layers.""" + for layer_wrapper in _find_kt_wrappers(model) or []: + moe = _profile_moe(layer_wrapper) + if moe is not None and hasattr(moe, "reset_profile_stats"): + moe.reset_profile_stats() + + +def _split_timer_key(key: str) -> tuple[str, str] | None: + suffix = ".total_ns" + if not key.endswith(suffix): + return None + timer = key[: -len(suffix)] + if timer.startswith("wrapper."): + return "wrapper", timer[len("wrapper.") :] + if timer.startswith("tp."): + parts = timer.split(".", 2) + if len(parts) == 3: + return f"tp.{parts[1]}", parts[2] + return None + + +def _parent_stage(stage: str) -> str | None: + if stage == "backward.down.total" or stage == "backward.gate_up.total": + return "backward.total" + if stage.startswith("backward.down."): + return "backward.down.total" + if stage.startswith("backward.gate_up."): + return "backward.gate_up.total" + if stage.endswith(".total"): + return None + if stage.startswith("forward."): + return "forward.total" + if stage.startswith("backward."): + return "backward.total" + if stage.startswith("tp.forward."): + return "tp.forward.total" + if stage.startswith("tp.backward."): + return "tp.backward.total" + return None + + +def _aggregate_rows(profile: dict[str, Any]) -> list[dict[str, float | str]]: + totals: dict[tuple[str, str], dict[str, float]] = defaultdict( + lambda: {"total_ns": 0.0, "calls": 0.0, "tokens": 0.0} + ) + for raw in profile.get("layers", {}).values(): + scope_tokens: dict[str, float] = {"wrapper": raw.get("wrapper.tokens", 0.0)} + tp_count = int(raw.get("tp_count", 0.0)) + for tp_idx in range(tp_count): + scope_tokens[f"tp.{tp_idx}"] = raw.get(f"tp.{tp_idx}.tokens", 0.0) + + for key, total_ns in raw.items(): + parsed = _split_timer_key(key) + if parsed is None or total_ns <= 0.0: + continue + scope, stage = parsed + calls = raw.get(key[: -len("total_ns")] + "calls", 0.0) + row = totals[(scope, stage)] + row["total_ns"] += total_ns + row["calls"] += calls + row["tokens"] += scope_tokens.get(scope, 0.0) + + rows: list[dict[str, float | str]] = [] + for (scope, stage), values in totals.items(): + calls = values["calls"] + tokens = values["tokens"] + parent = _parent_stage(stage) + parent_ns = totals.get((scope, parent), {}).get("total_ns", 0.0) if parent else 0.0 + rows.append( + { + "scope": scope, + "stage": stage, + "calls": calls, + "total_ms": values["total_ns"] / 1e6, + "avg_ms": values["total_ns"] / calls / 1e6 if calls else 0.0, + "us_per_token": values["total_ns"] / tokens / 1e3 if tokens else 0.0, + "parent_pct": values["total_ns"] / parent_ns * 100.0 if parent_ns else 0.0, + } + ) + return sorted(rows, key=lambda row: (str(row["scope"]), -float(row["total_ms"]), str(row["stage"]))) + + +def format_kt_sft_profile(profile: dict[str, Any]) -> str: + """Format a compact cross-layer timing table while preserving TP scopes.""" + if not profile.get("enabled"): + return "KT SFT profiler disabled or no profiled KT layers found. Set KT_SFT_PROFILE=1 before model creation." + + rows = _aggregate_rows(profile) + header = ( + f"{'scope':<9} {'stage':<34} {'calls':>7} {'total_ms':>11} " + f"{'avg_ms':>10} {'us/token':>10} {'parent%':>9}" + ) + lines = [header, "-" * len(header)] + for row in rows: + lines.append( + f"{row['scope']:<9} {row['stage']:<34} {row['calls']:>7.0f} " + f"{row['total_ms']:>11.3f} {row['avg_ms']:>10.3f} " + f"{row['us_per_token']:>10.3f} {row['parent_pct']:>8.1f}%" + ) + return "\n".join(lines) diff --git a/kt-kernel/test/per_commit/test_sft_profiler.py b/kt-kernel/test/per_commit/test_sft_profiler.py new file mode 100644 index 000000000..fbe5935b8 --- /dev/null +++ b/kt-kernel/test/per_commit/test_sft_profiler.py @@ -0,0 +1,68 @@ +from types import SimpleNamespace + +from kt_kernel.sft.profiler import collect_kt_sft_profile, format_kt_sft_profile, reset_kt_sft_profile + + +class _FakeMoe: + def __init__(self): + self.reset_calls = 0 + self.get_reset_values = [] + + def get_profile_stats(self, reset=False): + self.get_reset_values.append(reset) + return { + "layer_idx": 3, + "tp_count": 1, + "wrapper.enabled": 1, + "wrapper.tokens": 8, + "wrapper.tp.forward.total.total_ns": 2_000_000, + "wrapper.tp.forward.total.calls": 2, + "wrapper.tp.forward.numa_compute.total_ns": 1_500_000, + "wrapper.tp.forward.numa_compute.calls": 2, + "tp.0.enabled": 1, + "tp.0.tokens": 8, + "tp.0.forward.total.total_ns": 1_400_000, + "tp.0.forward.total.calls": 2, + "tp.0.forward.route.total_ns": 140_000, + "tp.0.forward.route.calls": 2, + "tp.0.backward.total.total_ns": 5_000_000, + "tp.0.backward.total.calls": 1, + "tp.0.backward.down.total.total_ns": 2_000_000, + "tp.0.backward.down.total.calls": 1, + } + + def reset_profile_stats(self): + self.reset_calls += 1 + + +def _fake_model(moe): + backend = SimpleNamespace(moe=moe) + layer = SimpleNamespace(layer_idx=3, wrapper=backend) + return SimpleNamespace(_kt_wrappers=[layer]) + + +def test_collect_and_format_profile(): + moe = _FakeMoe() + profile = collect_kt_sft_profile(_fake_model(moe), reset=True) + + assert profile["enabled"] is True + assert list(profile["layers"]) == [3] + assert moe.get_reset_values == [True] + + output = format_kt_sft_profile(profile) + assert "tp.forward.numa_compute" in output + assert "forward.route" in output + assert "75.0%" in output + assert "10.0%" in output + assert "40.0%" in output + + +def test_reset_and_disabled_profile(): + moe = _FakeMoe() + model = _fake_model(moe) + reset_kt_sft_profile(model) + assert moe.reset_calls == 1 + + empty = collect_kt_sft_profile(SimpleNamespace(_kt_wrappers=[])) + assert empty == {"enabled": False, "layers": {}} + assert "disabled" in format_kt_sft_profile(empty) From c6f4211346e1aa7b42f66551c33f9a131a2e4dc3 Mon Sep 17 00:00:00 2001 From: yyj Date: Fri, 17 Jul 2026 11:29:41 +0800 Subject: [PATCH 11/20] [fix](kt-kernel): reuse inference BF16 kernel for SFT --- kt-kernel/ext_bindings.cpp | 5 +- kt-kernel/operators/amx/bf16-moe.hpp | 2 + .../operators/amx/la/amx_raw_buffers.hpp | 108 +++++++++++++++++- kt-kernel/operators/amx/sft_moe.hpp | 18 ++- .../amx/test/test_raw_bf16_repack.cpp | 98 ++++++++++++++++ 5 files changed, 218 insertions(+), 13 deletions(-) create mode 100644 kt-kernel/operators/amx/test/test_raw_bf16_repack.cpp diff --git a/kt-kernel/ext_bindings.cpp b/kt-kernel/ext_bindings.cpp index f3408f935..9ce751752 100644 --- a/kt-kernel/ext_bindings.cpp +++ b/kt-kernel/ext_bindings.cpp @@ -819,7 +819,7 @@ PYBIND11_MODULE(kt_kernel_ext, m) { #endif #if defined(__AVX512BF16__) // SFT MoE with LoRA support (BF16, INT8, INT4, AWQ, K2) - bind_moe_sft_module>(moe_module, "AMXBF16_SFT_MOE"); + bind_moe_sft_module>(moe_module, "AMXBF16_SFT_MOE"); bind_moe_sft_module>(moe_module, "AMXInt8_SFT_MOE"); bind_moe_sft_module>(moe_module, "AMXInt4_SFT_MOE"); // bind_moe_sft_module>(moe_module, "AMXInt4_1_SFT_MOE"); @@ -828,7 +828,8 @@ PYBIND11_MODULE(kt_kernel_ext, m) { // bind_moe_sft_module>(moe_module, // "AMXInt4_KGroup_SFT_MOE"); // SFT MoE with SkipLoRA=true (skip all LoRA computation in backward, only compute base weight grad_input) - bind_moe_sft_module>(moe_module, "AMXBF16_SFT_MOE_SkipLoRA"); + bind_moe_sft_module>(moe_module, + "AMXBF16_SFT_MOE_SkipLoRA"); bind_moe_sft_module>(moe_module, "AMXInt8_SFT_MOE_SkipLoRA"); bind_moe_sft_module>(moe_module, "AMXInt4_SFT_MOE_SkipLoRA"); // bind_moe_sft_module>(moe_module, diff --git a/kt-kernel/operators/amx/bf16-moe.hpp b/kt-kernel/operators/amx/bf16-moe.hpp index 389446ed6..375aaef9d 100644 --- a/kt-kernel/operators/amx/bf16-moe.hpp +++ b/kt-kernel/operators/amx/bf16-moe.hpp @@ -30,6 +30,8 @@ template class AMX_BF16_MOE_TP : public AMX_MOE_BASE> { using Base = AMX_MOE_BASE>; + + protected: using Base::config_; using Base::down_ba_; using Base::down_bb_; diff --git a/kt-kernel/operators/amx/la/amx_raw_buffers.hpp b/kt-kernel/operators/amx/la/amx_raw_buffers.hpp index 86fee938a..8b28b411e 100644 --- a/kt-kernel/operators/amx/la/amx_raw_buffers.hpp +++ b/kt-kernel/operators/amx/la/amx_raw_buffers.hpp @@ -133,16 +133,13 @@ struct BufferBBF16Impl { } void set_data(void* new_ptr) { b = reinterpret_cast(new_ptr); } - void from_mat(ggml_bf16_t* src, int ith, int nth) { - auto [n_start, n_end] = K::split_range_n(n, ith, nth); - int n_block_begin = n_start; - int n_block_size = n_end - n_block_begin; + void pack_block(ggml_bf16_t* src, int src_stride, int n_block_begin, int n_block_size) { for (int n_begin = 0; n_begin < n_block_size; n_begin += N_STEP) { for (int k_block_begin = 0; k_block_begin < k; k_block_begin += K_BLOCK) { int k_block_size = std::min(K_BLOCK, k - k_block_begin); for (int k_begin = 0; k_begin < k_block_size; k_begin += K_STEP) { for (int i = 0; i < N_STEP; i++) { - __m512i* s = (__m512i*)(src + (n_block_begin + n_begin + i) * k + k_block_begin + k_begin); + __m512i* s = (__m512i*)(src + (n_begin + i) * src_stride + k_block_begin + k_begin); __m512i* d = (__m512i*)(b + n_block_begin * k + k_block_begin * n_block_size + n_begin * k_block_size + k_begin * N_STEP + i * K_STEP); avx512_copy_32xbf16(s, d); @@ -155,6 +152,107 @@ struct BufferBBF16Impl { } } } + + void from_mat(ggml_bf16_t* src, int ith, int nth) { + auto [n_start, n_end] = K::split_range_n(n, ith, nth); + int n_block_begin = n_start; + int n_block_size = n_end - n_block_begin; + pack_block(src + n_block_begin * k, k, n_block_begin, n_block_size); + } + + void from_mat_transposed(ggml_bf16_t* src, int src_n, int src_k, int ith, int nth) { + assert(n == src_k && k == src_n); + auto [n_start, n_end] = K::split_range_n(n, ith, nth); + int n_block_begin = n_start; + int n_block_size = n_end - n_block_begin; + if (n_block_size <= 0) return; + + thread_local std::vector strip; + strip.resize((size_t)n_block_size * k); + constexpr int TILE = 32; + for (int c_tile = 0; c_tile < k; c_tile += TILE) { + int c_end = std::min(c_tile + TILE, k); + for (int r_tile = 0; r_tile < n_block_size; r_tile += TILE) { + int r_end = std::min(r_tile + TILE, n_block_size); + for (int c = c_tile; c < c_end; c++) { + for (int r = r_tile; r < r_end; r++) { + strip[(size_t)r * k + c] = src[(size_t)c * src_k + n_block_begin + r]; + } + } + } + } + pack_block(strip.data(), k, n_block_begin, n_block_size); + } + + void to_mat(ggml_bf16_t* dst, int ith, int nth) const { + auto [n_start, n_end] = K::split_range_n(n, ith, nth); + int n_block_begin = n_start; + int n_block_size = n_end - n_block_begin; + if (n_block_size <= 0) return; + + alignas(64) ggml_bf16_t tile_copy[N_STEP * K_STEP]; + for (int n_begin = 0; n_begin < n_block_size; n_begin += N_STEP) { + for (int k_block_begin = 0; k_block_begin < k; k_block_begin += K_BLOCK) { + int k_block_size = std::min(K_BLOCK, k - k_block_begin); + for (int k_begin = 0; k_begin < k_block_size; k_begin += K_STEP) { + const ggml_bf16_t* tile_src = + b + n_block_begin * k + k_block_begin * n_block_size + n_begin * k_block_size + k_begin * N_STEP; + memcpy(tile_copy, tile_src, sizeof(tile_copy)); + transpose_16x16_32bit((__m512i*)tile_copy); + transpose_16x16_32bit((__m512i*)(tile_copy + TILE_N * K_STEP)); + for (int i = 0; i < N_STEP; i++) { + __m512i* s = (__m512i*)(tile_copy + i * K_STEP); + __m512i* d = (__m512i*)(dst + (size_t)(n_block_begin + n_begin + i) * k + k_block_begin + k_begin); + avx512_copy_32xbf16(s, d); + } + } + } + } + } + + void from_bb_transposed(const BufferBBF16Impl& src, int ith, int nth) { + assert(n == src.k && k == src.n); + auto [n_start, n_end] = K::split_range_n(n, ith, nth); + int dst_n_block_begin = n_start; + int dst_n_block_size = n_end - dst_n_block_begin; + if (dst_n_block_size <= 0) return; + + auto tile_ptr = [](ggml_bf16_t* base, int total_n, int total_k, int abs_n, int abs_k) { + int n_block_begin = abs_n / N_BLOCK * N_BLOCK; + int n_within = abs_n - n_block_begin; + int n_block_size = std::min(N_BLOCK, total_n - n_block_begin); + int k_block_begin = abs_k / K_BLOCK * K_BLOCK; + int k_within = abs_k - k_block_begin; + return base + (size_t)n_block_begin * total_k + (size_t)k_block_begin * n_block_size + + (size_t)n_within * std::min(K_BLOCK, total_k - k_block_begin) + (size_t)k_within * N_STEP; + }; + + alignas(64) ggml_bf16_t src_tile[N_STEP * K_STEP]; + alignas(64) ggml_bf16_t dst_tile[N_STEP * K_STEP]; + for (int dst_n = 0; dst_n < dst_n_block_size; dst_n += N_STEP) { + for (int dst_k_block = 0; dst_k_block < k; dst_k_block += K_BLOCK) { + int dst_k_block_size = std::min(K_BLOCK, k - dst_k_block); + for (int dst_k = 0; dst_k < dst_k_block_size; dst_k += K_STEP) { + int abs_dst_n = dst_n_block_begin + dst_n; + int abs_dst_k = dst_k_block + dst_k; + ggml_bf16_t* src_ptr = tile_ptr(src.b, src.n, src.k, abs_dst_k, abs_dst_n); + memcpy(src_tile, src_ptr, sizeof(src_tile)); + transpose_16x16_32bit((__m512i*)src_tile); + transpose_16x16_32bit((__m512i*)(src_tile + TILE_N * K_STEP)); + + for (int i = 0; i < N_STEP; i++) { + for (int j = 0; j < K_STEP; j++) { + dst_tile[j * K_STEP + i] = src_tile[i * K_STEP + j]; + } + } + transpose_16x16_32bit((__m512i*)dst_tile); + transpose_16x16_32bit((__m512i*)(dst_tile + TILE_N * K_STEP)); + ggml_bf16_t* dst_ptr = tile_ptr(b, n, k, abs_dst_n, abs_dst_k); + memcpy(dst_ptr, dst_tile, sizeof(dst_tile)); + } + } + } + } ggml_bf16_t* get_submat(int n, int k, int n_begin, int k_begin) { int n_block_begin = n_begin / N_BLOCK * N_BLOCK; n_begin -= n_block_begin; diff --git a/kt-kernel/operators/amx/sft_moe.hpp b/kt-kernel/operators/amx/sft_moe.hpp index d0deb24b5..a9ba049f7 100644 --- a/kt-kernel/operators/amx/sft_moe.hpp +++ b/kt-kernel/operators/amx/sft_moe.hpp @@ -30,6 +30,7 @@ #include "../sft_profile.hpp" #include "ggml.h" #include "la/amx_kernels.hpp" +#include "la/amx_raw_kernels.hpp" #include "la/avx_kernels.hpp" #include "moe.hpp" @@ -220,6 +221,8 @@ struct supports_standard_mat_mul : std::false_type {}; template <> struct supports_standard_mat_mul : std::true_type {}; template <> +struct supports_standard_mat_mul : std::true_type {}; +template <> struct supports_standard_mat_mul : std::true_type {}; template <> struct supports_standard_mat_mul : std::true_type {}; @@ -238,6 +241,8 @@ struct has_bb_transposed_repack : std::false_type {}; template <> struct has_bb_transposed_repack : std::true_type {}; template <> +struct has_bb_transposed_repack : std::true_type {}; +template <> struct has_bb_transposed_repack : std::true_type {}; template inline constexpr bool has_bb_transposed_repack_v = has_bb_transposed_repack::value; @@ -2002,10 +2007,11 @@ class AMX_SFT_MOE_TP : public BaseMOE { token_offset += cache.m_local_num_cache[expert_idx]; } - if constexpr (std::is_same_v && amx::AMX_AVAILABLE) { - constexpr int TILE_M = amx::GemmKernel224BF::M_STEP; - constexpr int TILE_N = amx::GemmKernel224BF::N_STEP; - constexpr int TILE_K = amx::GemmKernel224BF::K_STEP; + if constexpr ((std::is_same_v || std::is_same_v) && + amx::AMX_AVAILABLE) { + constexpr int TILE_M = T::M_STEP; + constexpr int TILE_N = T::N_STEP; + constexpr int TILE_K = T::K_STEP; static_assert(TILE_M == 32 && TILE_N == 32 && TILE_K == 32, "base-weight gradient tile packing assumes 32x32x32 AMX BF16 tiles"); @@ -2068,7 +2074,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { } } amx::transpose_16x16_32bit(reinterpret_cast<__m512i*>(b_tile)); - amx::transpose_16x16_32bit(reinterpret_cast<__m512i*>(b_tile + amx::GemmKernel224BF::TILE_N * TILE_K)); + amx::transpose_16x16_32bit(reinterpret_cast<__m512i*>(b_tile + T::TILE_N * TILE_K)); } ggml_bf16_t* down_dst = gdp + (size_t)expert_idx * H * F; @@ -2139,7 +2145,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { } } amx::transpose_16x16_32bit(reinterpret_cast<__m512i*>(b_tile)); - amx::transpose_16x16_32bit(reinterpret_cast<__m512i*>(b_tile + amx::GemmKernel224BF::TILE_N * TILE_K)); + amx::transpose_16x16_32bit(reinterpret_cast<__m512i*>(b_tile + T::TILE_N * TILE_K)); } // Gate and up each consume all four C tiles, so retain C across K in two separate passes. diff --git a/kt-kernel/operators/amx/test/test_raw_bf16_repack.cpp b/kt-kernel/operators/amx/test/test_raw_bf16_repack.cpp new file mode 100644 index 000000000..50bf9853d --- /dev/null +++ b/kt-kernel/operators/amx/test/test_raw_bf16_repack.cpp @@ -0,0 +1,98 @@ +#include +#include +#include +#include +#include + +#include "../la/amx_kernels.hpp" +#include "../la/amx_raw_kernels.hpp" + +namespace { + +using Kernel = amx::GemmKernel224BF16; +using BufferB = Kernel::BufferB; + +void* alloc_buffer(size_t bytes) { + void* ptr = std::aligned_alloc(64, (bytes + 63) / 64 * 64); + if (ptr == nullptr) std::abort(); + std::memset(ptr, 0, bytes); + return ptr; +} + +void fill_random(std::vector& values, unsigned seed) { + std::mt19937 generator(seed); + std::uniform_real_distribution distribution(-1.0f, 1.0f); + for (auto& value : values) value = GGML_FP32_TO_BF16(distribution(generator)); +} + +void from_mat(BufferB& buffer, ggml_bf16_t* source) { + int nth = Kernel::recommended_nth(buffer.n); + for (int ith = 0; ith < nth; ith++) buffer.from_mat(source, ith, nth); +} + +void to_mat(const BufferB& buffer, ggml_bf16_t* destination) { + int nth = Kernel::recommended_nth(buffer.n); + for (int ith = 0; ith < nth; ith++) buffer.to_mat(destination, ith, nth); +} + +void from_mat_transposed(BufferB& buffer, ggml_bf16_t* source, int source_n, int source_k) { + int nth = Kernel::recommended_nth(buffer.n); + for (int ith = 0; ith < nth; ith++) buffer.from_mat_transposed(source, source_n, source_k, ith, nth); +} + +void from_bb_transposed(BufferB& destination, const BufferB& source) { + int nth = Kernel::recommended_nth(destination.n); + for (int ith = 0; ith < nth; ith++) destination.from_bb_transposed(source, ith, nth); +} + +bool run_case(int n, int k) { + const size_t count = (size_t)n * k; + std::vector source(count); + std::vector roundtrip(count); + std::vector transposed(count); + std::vector direct_transposed(count); + fill_random(source, (unsigned)(n * 31 + k)); + + void* forward_memory = alloc_buffer(BufferB::required_size(n, k)); + void* expected_memory = alloc_buffer(BufferB::required_size(k, n)); + void* direct_memory = alloc_buffer(BufferB::required_size(k, n)); + BufferB forward(n, k, forward_memory); + BufferB expected(k, n, expected_memory); + BufferB direct(k, n, direct_memory); + + from_mat(forward, source.data()); + to_mat(forward, roundtrip.data()); + from_mat_transposed(expected, source.data(), n, k); + from_bb_transposed(direct, forward); + to_mat(expected, transposed.data()); + to_mat(direct, direct_transposed.data()); + + bool roundtrip_ok = std::memcmp(source.data(), roundtrip.data(), count * sizeof(ggml_bf16_t)) == 0; + bool direct_ok = std::memcmp(transposed.data(), direct_transposed.data(), count * sizeof(ggml_bf16_t)) == 0; + bool transpose_ok = true; + for (int row = 0; row < n && transpose_ok; row++) { + for (int column = 0; column < k; column++) { + if (source[(size_t)row * k + column].bits != transposed[(size_t)column * n + row].bits) { + transpose_ok = false; + break; + } + } + } + + std::free(forward_memory); + std::free(expected_memory); + std::free(direct_memory); + std::printf("raw BF16 repack %dx%d: roundtrip=%s transpose=%s direct=%s\n", n, k, roundtrip_ok ? "PASS" : "FAIL", + transpose_ok ? "PASS" : "FAIL", direct_ok ? "PASS" : "FAIL"); + return roundtrip_ok && transpose_ok && direct_ok; +} + +} // namespace + +int main() { + bool passed = true; + passed = run_case(64, 64) && passed; + passed = run_case(768, 2048) && passed; + passed = run_case(2048, 768) && passed; + return passed ? 0 : 1; +} From 109b403b633865f75c1a86411f3c2749c229725d Mon Sep 17 00:00:00 2001 From: yyj Date: Fri, 17 Jul 2026 12:04:36 +0800 Subject: [PATCH 12/20] [perf](kt-kernel): add fine-grained Full-FT profiling --- kt-kernel/operators/amx/sft_moe.hpp | 20 +++++++++-- kt-kernel/operators/moe-sft-tp.hpp | 8 +++++ kt-kernel/operators/sft_profile.hpp | 24 +++++++++++++ kt-kernel/python/sft/autograd.py | 33 +++++++++++------- kt-kernel/python/sft/layer.py | 54 ++++++++++++++++------------- kt-kernel/python/sft/profiler.py | 4 +++ 6 files changed, 103 insertions(+), 40 deletions(-) diff --git a/kt-kernel/operators/amx/sft_moe.hpp b/kt-kernel/operators/amx/sft_moe.hpp index a9ba049f7..451732ad1 100644 --- a/kt-kernel/operators/amx/sft_moe.hpp +++ b/kt-kernel/operators/amx/sft_moe.hpp @@ -949,6 +949,8 @@ class AMX_SFT_MOE_TP : public BaseMOE { void forward_sft(int qlen, int k, const int64_t* expert_ids, const float* weights, const void* input, void* output, bool save_for_backward) { SFTProfileScope total_scope(profiler_, SFTProfileStage::FwdTotal); + SFTProfileScope checkpoint_scope( + profiler_, save_for_backward ? SFTProfileStage::FwdRecomputeTotal : SFTProfileStage::FwdInitialTotal); auto stage_start = profiler_.start(); SFT_POOL_LOG("fwd_enter", config_.layer_idx, tp_part_idx, qlen, cache_stack_top_, forward_pool_bytes_, @@ -1986,6 +1988,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { */ void backward_base_weight_grad(const ForwardCache& cache, int full_intermediate_size, void* grad_gate_proj, void* grad_up_proj, void* grad_down_proj) { + auto stage_start = profiler_.start(); const int H = config_.hidden_size; const int I = config_.intermediate_size; const int F = full_intermediate_size; @@ -2006,6 +2009,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { int expert_idx = cache.m_expert_id_map_cache[task_id]; token_offset += cache.m_local_num_cache[expert_idx]; } + profiler_.record(SFTProfileStage::BwdBaseWeightGradOffsets, stage_start); if constexpr ((std::is_same_v || std::is_same_v) && amx::AMX_AVAILABLE) { @@ -2022,6 +2026,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { const int total_tasks = activated_expert * tasks_per_expert; auto pool = config_.pool->get_subpool(tp_part_idx); + stage_start = profiler_.start(); pool->do_work_stealing_job( total_tasks, [](int _) { T::config(); }, [&, i_tiles, h_tiles, tasks_per_expert](int task_id) { @@ -2174,6 +2179,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { } }, nullptr); + profiler_.record(SFTProfileStage::BwdBaseWeightGradAmx, stage_start); return; } @@ -2189,16 +2195,17 @@ class AMX_SFT_MOE_TP : public BaseMOE { float* acc_up = acc_gate + (size_t)I * H; // [I, H] float* acc_down = acc_up + (size_t)I * H; // [H, I] + stage_start = profiler_.start(); std::memset(acc_gate, 0, (size_t)I * H * sizeof(float)); std::memset(acc_up, 0, (size_t)I * H * sizeof(float)); std::memset(acc_down, 0, (size_t)H * I * sizeof(float)); + profiler_.record(SFTProfileStage::BwdBaseWeightGradZero, stage_start); + stage_start = profiler_.start(); for (int t = 0; t < m; t++) { const ggml_bf16_t* input_row = m_local_input_ptr_[expert_idx] + (size_t)t * H; const ggml_bf16_t* gate_grad_row = grad_gate_output_ + (size_t)(pos_start + t) * I; const ggml_bf16_t* up_grad_row = grad_up_output_ + (size_t)(pos_start + t) * I; - const ggml_bf16_t* inter_row = cache.intermediate_cache + (size_t)(pos_start + t) * I; - const ggml_bf16_t* grad_out_row = base_grad_output_bf16_ptr_[expert_idx] + (size_t)t * H; // gate_proj grad: [I, H] += grad_gate_out[t]^T @ input[t] for (int i = 0; i < I; i++) { @@ -2210,7 +2217,13 @@ class AMX_SFT_MOE_TP : public BaseMOE { acc_up[i * H + h] += gu * inp; } } + } + profiler_.record(SFTProfileStage::BwdBaseWeightGradGateUp, stage_start); + stage_start = profiler_.start(); + for (int t = 0; t < m; t++) { + const ggml_bf16_t* inter_row = cache.intermediate_cache + (size_t)(pos_start + t) * I; + const ggml_bf16_t* grad_out_row = base_grad_output_bf16_ptr_[expert_idx] + (size_t)t * H; // down_proj grad: [H, I] += grad_output[t]^T @ intermediate[t] for (int h = 0; h < H; h++) { float go = GGML_BF16_TO_FP32(grad_out_row[h]); @@ -2219,8 +2232,10 @@ class AMX_SFT_MOE_TP : public BaseMOE { } } } + profiler_.record(SFTProfileStage::BwdBaseWeightGradDown, stage_start); // Convert FP32 accumulators to BF16 and store + stage_start = profiler_.start(); for (int i = 0; i < I; i++) { for (int h = 0; h < H; h++) { ggp[(size_t)expert_idx * F * H + (size_t)i * H + h] = GGML_FP32_TO_BF16(acc_gate[i * H + h]); @@ -2232,6 +2247,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { gdp[(size_t)expert_idx * H * F + (size_t)h * F + i] = GGML_FP32_TO_BF16(acc_down[h * I + i]); } } + profiler_.record(SFTProfileStage::BwdBaseWeightGradStore, stage_start); } } diff --git a/kt-kernel/operators/moe-sft-tp.hpp b/kt-kernel/operators/moe-sft-tp.hpp index 79ecadf49..4605d6c7b 100644 --- a/kt-kernel/operators/moe-sft-tp.hpp +++ b/kt-kernel/operators/moe-sft-tp.hpp @@ -370,6 +370,7 @@ class TP_MOE_SFT : public TP_MOE { } } else if (config.gate_proj != nullptr) { // printf("TP_MOE_SFT: From BF16 with partitioning\n"); + auto reload_stage_start = profiler_.start(); // Temporary storage for partitioned weights std::vector temp_gate(tp_count); @@ -414,27 +415,34 @@ class TP_MOE_SFT : public TP_MOE { }, nullptr); } + profiler_.record(SFTProfileStage::BaseWeightReloadPartition, reload_stage_start); // Step 2: Set weight pointers BEFORE load_weights (Bug #24 fix) + reload_stage_start = profiler_.start(); for (int i = 0; i < tp_count; i++) { tps[i]->set_physical_to_logical_map(config.physical_to_logical_map); tps[i]->set_weight_pointers_for_forward(temp_gate[i], temp_up[i], temp_down[i]); } pool->dispense_backend()->do_numa_job([this](int numa_id) { tps[numa_id]->load_weights(); }); + profiler_.record(SFTProfileStage::BaseWeightReloadForwardPack, reload_stage_start); // Step 3: Prepare backward weights (this also clears weight pointers) + reload_stage_start = profiler_.start(); for (int i = 0; i < tp_count; i++) { if (!config.share_backward_bb) { tps[i]->prepare_bwd(temp_gate[i], temp_up[i], temp_down[i]); } } + profiler_.record(SFTProfileStage::BaseWeightReloadBackwardPack, reload_stage_start); + reload_stage_start = profiler_.start(); for (int i = 0; i < tp_count; i++) { delete[] (temp_gate[i]); delete[] (temp_up[i]); delete[] (temp_down[i]); } + profiler_.record(SFTProfileStage::BaseWeightReloadCleanup, reload_stage_start); } else { // Other loading methods (from loader or file) for (int i = 0; i < tp_count; i++) { diff --git a/kt-kernel/operators/sft_profile.hpp b/kt-kernel/operators/sft_profile.hpp index 17dc77198..91ddefb77 100644 --- a/kt-kernel/operators/sft_profile.hpp +++ b/kt-kernel/operators/sft_profile.hpp @@ -16,6 +16,8 @@ enum class SFTProfileStage : uint8_t { // NUMA-local forward stages. FwdTotal, + FwdInitialTotal, + FwdRecomputeTotal, FwdSetup, FwdRoute, FwdBufferSetup, @@ -48,6 +50,12 @@ enum class SFTProfileStage : uint8_t { BwdGateUpLora, BwdRouterGrad, BwdBaseWeightGrad, + BwdBaseWeightGradOffsets, + BwdBaseWeightGradAmx, + BwdBaseWeightGradZero, + BwdBaseWeightGradGateUp, + BwdBaseWeightGradDown, + BwdBaseWeightGradStore, // TP wrapper and weight-layout stages. TpFwdTotal, @@ -61,12 +69,18 @@ enum class SFTProfileStage : uint8_t { TpBwdRouterGradMerge, BackwardRepack, BaseWeightReload, + BaseWeightReloadPartition, + BaseWeightReloadForwardPack, + BaseWeightReloadBackwardPack, + BaseWeightReloadCleanup, Count, }; inline constexpr std::array(SFTProfileStage::Count)> kSFTProfileStageNames = { "forward.total", + "forward.initial_total", + "forward.recompute_total", "forward.setup", "forward.route", "forward.buffer_setup", @@ -97,6 +111,12 @@ inline constexpr std::array(SFTProfileStage::Co "backward.gate_up.lora", "backward.router_grad", "backward.base_weight_grad", + "backward.base_weight_grad.offsets", + "backward.base_weight_grad.amx", + "backward.base_weight_grad.zero", + "backward.base_weight_grad.gate_up", + "backward.base_weight_grad.down", + "backward.base_weight_grad.store", "tp.forward.total", "tp.forward.numa_compute", "tp.forward.merge", @@ -108,6 +128,10 @@ inline constexpr std::array(SFTProfileStage::Co "tp.backward.router_grad_merge", "weights.backward_repack", "weights.base_reload", + "weights.base_reload.partition", + "weights.base_reload.forward_pack", + "weights.base_reload.backward_pack", + "weights.base_reload.cleanup", }; inline bool sft_profile_enabled_from_env() { diff --git a/kt-kernel/python/sft/autograd.py b/kt-kernel/python/sft/autograd.py index 98c92ed97..bb9ce7974 100644 --- a/kt-kernel/python/sft/autograd.py +++ b/kt-kernel/python/sft/autograd.py @@ -78,7 +78,8 @@ def forward( # Rank 0: sync CPU result and split by real lengths if rank == 0: - cpu_output = wrapper.sync_forward(output_device=original_device) + with torch.profiler.record_function("kt.sft.cpu_forward_sync"): + cpu_output = wrapper.sync_forward(output_device=original_device) cpu_output = cpu_output.to(dtype=original_dtype).view(total_qlen, hidden_size) offsets = _qlen_offsets(all_qlens_list) scatter_list = [cpu_output[offsets[i] : offsets[i + 1]].contiguous() for i in range(world_size)] @@ -98,7 +99,8 @@ def forward( del output_flat elif wrapper is not None: # Single-GPU: sync directly - cpu_output = wrapper.sync_forward(output_device=original_device) + with torch.profiler.record_function("kt.sft.cpu_forward_sync"): + cpu_output = wrapper.sync_forward(output_device=original_device) output = cpu_output.view(batch_size, seq_len, hidden_size).to(dtype=original_dtype) else: # Broadcast-only rank (no wrapper) @@ -141,12 +143,14 @@ def forward( def backward(ctx, grad_output: torch.Tensor): # Wait for any in-flight async repack before recompute forward uses the pool if getattr(ctx.wrapper, "share_backward_bb", False): - ctx.wrapper.wait_backward_repack() + with torch.profiler.record_function("kt.sft.wait_backward_repack"): + ctx.wrapper.wait_backward_repack() # Access saved_tensors FIRST — under non-reentrant checkpoint this # triggers the unpack hook which runs a full decoder-layer recompute, # populating the C++ cache before we call wrapper.backward(). - _ = ctx.saved_tensors + with torch.profiler.record_function("kt.sft.checkpoint_recompute"): + _ = ctx.saved_tensors qlen = ctx.qlen hidden_size = ctx.hidden_size @@ -191,10 +195,11 @@ def backward(ctx, grad_output: torch.Tensor): all_go = torch.cat(gathered_go, dim=0) total_qlen = int(all_go.shape[0]) - backward_out = ctx.wrapper.backward( - all_go, - output_device=ctx.original_device, - ) + with torch.profiler.record_function("kt.sft.cpu_backward"): + backward_out = ctx.wrapper.backward( + all_go, + output_device=ctx.original_device, + ) if isinstance(backward_out, tuple) and len(backward_out) == 2: all_grad_input, all_grad_weights = backward_out elif isinstance(backward_out, tuple) and len(backward_out) == 3: @@ -236,10 +241,11 @@ def backward(ctx, grad_output: torch.Tensor): elif not ctx.use_broadcast: # ---- Single-GPU path ---- grad_output_flat = grad_output.view(qlen, hidden_size) - backward_out = ctx.wrapper.backward( - grad_output_flat, - output_device=ctx.original_device, - ) + with torch.profiler.record_function("kt.sft.cpu_backward"): + backward_out = ctx.wrapper.backward( + grad_output_flat, + output_device=ctx.original_device, + ) ctx.wrapper._kt_has_cached_forward = False if isinstance(backward_out, tuple) and len(backward_out) == 2: grad_input, grad_weights = backward_out @@ -259,7 +265,8 @@ def backward(ctx, grad_output: torch.Tensor): # Trigger async repack for next MoE layer in backward order next_bwd = getattr(ctx.wrapper, "_next_backward_wrapper", None) if next_bwd is not None and getattr(next_bwd, "share_backward_bb", False): - next_bwd.submit_backward_repack() + with torch.profiler.record_function("kt.sft.submit_backward_repack"): + next_bwd.submit_backward_repack() # Base weight gradients: return C++-written grad buffers in full mode, None otherwise if ctx.full_weight_grad and ctx.wrapper is not None: diff --git a/kt-kernel/python/sft/layer.py b/kt-kernel/python/sft/layer.py index fc12949d4..7db6398b0 100644 --- a/kt-kernel/python/sft/layer.py +++ b/kt-kernel/python/sft/layer.py @@ -136,7 +136,8 @@ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: # Check if we need to use distributed broadcast (only rank 0 has KT kernel) use_broadcast = dist_on and self.wrapper is None - topk_ids, topk_weights = self._compute_routing(hidden_states) + with torch.profiler.record_function("kt.sft.routing"): + topk_ids, topk_weights = self._compute_routing(hidden_states) train_lora = self._peft_lora_modules is not None and len(self._peft_lora_modules) > 0 full_weight_grad = self._full_weight_grad @@ -157,15 +158,17 @@ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: # In full_weight_grad mode, sync base weights after optimizer step if full_weight_grad and getattr(self.wrapper, "_base_weights_dirty", False): - self.wrapper.update_base_weights() + with torch.profiler.record_function("kt.sft.base_weight_reload"): + self.wrapper.update_base_weights() self.wrapper._base_weights_dirty = False - gpu_output, all_qlens = self._submit_and_compute_gpu( - hidden_states, - topk_ids, - topk_weights, - save_for_backward_submit, - ) + with torch.profiler.record_function("kt.sft.submit_and_gpu_experts"): + gpu_output, all_qlens = self._submit_and_compute_gpu( + hidden_states, + topk_ids, + topk_weights, + save_for_backward_submit, + ) # Use KTMoEFunction whenever backward is needed so KT backward and LoRA # gradient paths remain connected. @@ -184,23 +187,24 @@ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: if self.wrapper.gate_proj_buf is not None: lora_ref = self.wrapper.gate_proj_buf - moe_output = KTMoEFunction.apply( - hidden_states, - topk_ids, - topk_weights, - self.wrapper, - lora_ref, - self.hidden_size, - self.moe_config.num_experts_per_tok, - self.layer_idx, - save_for_backward, - train_lora, - all_qlens, - # Base weight params for full mode gradient flow - self.wrapper.gate_proj_buf if full_weight_grad and self.wrapper is not None else None, - self.wrapper.up_proj_buf if full_weight_grad and self.wrapper is not None else None, - self.wrapper.down_proj_buf if full_weight_grad and self.wrapper is not None else None, - ) + with torch.profiler.record_function("kt.sft.autograd_apply_and_cpu_sync"): + moe_output = KTMoEFunction.apply( + hidden_states, + topk_ids, + topk_weights, + self.wrapper, + lora_ref, + self.hidden_size, + self.moe_config.num_experts_per_tok, + self.layer_idx, + save_for_backward, + train_lora, + all_qlens, + # Base weight params for full mode gradient flow + self.wrapper.gate_proj_buf if full_weight_grad and self.wrapper is not None else None, + self.wrapper.up_proj_buf if full_weight_grad and self.wrapper is not None else None, + self.wrapper.down_proj_buf if full_weight_grad and self.wrapper is not None else None, + ) else: moe_output = self._sync_forward_output_no_autograd( hidden_states=hidden_states, diff --git a/kt-kernel/python/sft/profiler.py b/kt-kernel/python/sft/profiler.py index 1cab47d26..6d1bb8336 100644 --- a/kt-kernel/python/sft/profiler.py +++ b/kt-kernel/python/sft/profiler.py @@ -63,6 +63,10 @@ def _split_timer_key(key: str) -> tuple[str, str] | None: def _parent_stage(stage: str) -> str | None: + if stage.startswith("backward.base_weight_grad."): + return "backward.base_weight_grad" + if stage.startswith("weights.base_reload."): + return "weights.base_reload" if stage == "backward.down.total" or stage == "backward.gate_up.total": return "backward.total" if stage.startswith("backward.down."): From 34d2102082c69fdbcb49e97fa653ce4352a2cd18 Mon Sep 17 00:00:00 2001 From: yyj Date: Fri, 17 Jul 2026 14:52:27 +0800 Subject: [PATCH 13/20] [perf](kt-kernel): batch BF16 Full-FT weight gradients Use one expert-aggregated tile driver for AVX512-BF16 and AMX base-weight gradients, and pack updated full-precision weights directly into TP BufferB layouts without temporary partitions. Add worker-local profiling and focused dWeight/strided-repack coverage. --- kt-kernel/CMakeLists.txt | 1 + .../operators/amx/la/amx_raw_buffers.hpp | 9 + kt-kernel/operators/amx/la/bf16_dweight.hpp | 168 +++++++++++++++++ kt-kernel/operators/amx/sft_moe.hpp | 176 +++++++++++++++++- .../operators/amx/test/test_bf16_dweight.cpp | 106 +++++++++++ .../amx/test/test_raw_bf16_repack.cpp | 35 ++++ kt-kernel/operators/moe-sft-tp.hpp | 164 +++++++++------- kt-kernel/operators/sft_profile.hpp | 29 ++- kt-kernel/python/sft/profiler.py | 2 + 9 files changed, 614 insertions(+), 76 deletions(-) create mode 100644 kt-kernel/operators/amx/la/bf16_dweight.hpp create mode 100644 kt-kernel/operators/amx/test/test_bf16_dweight.cpp diff --git a/kt-kernel/CMakeLists.txt b/kt-kernel/CMakeLists.txt index 1e7558b8f..f87c03df7 100644 --- a/kt-kernel/CMakeLists.txt +++ b/kt-kernel/CMakeLists.txt @@ -395,6 +395,7 @@ if(HOST_IS_X86) # 获取不带扩展名的文件名作为 target 名 get_filename_component(test_name ${test_src} NAME_WE) add_executable(${test_name} ${test_src} ${CMAKE_CURRENT_SOURCE_DIR}/cpu_backend/shared_mem_buffer.cpp) + target_compile_options(${test_name} PRIVATE ${ARCH_FLAGS}) target_link_libraries(${test_name} llama OpenMP::OpenMP_CXX numa) endforeach() endif() diff --git a/kt-kernel/operators/amx/la/amx_raw_buffers.hpp b/kt-kernel/operators/amx/la/amx_raw_buffers.hpp index 8b28b411e..fc3b67194 100644 --- a/kt-kernel/operators/amx/la/amx_raw_buffers.hpp +++ b/kt-kernel/operators/amx/la/amx_raw_buffers.hpp @@ -160,6 +160,15 @@ struct BufferBBF16Impl { pack_block(src + n_block_begin * k, k, n_block_begin, n_block_size); } + void from_mat_strided(ggml_bf16_t* src, int src_stride, int ith, int nth) { + assert(src_stride >= k); + auto [n_start, n_end] = K::split_range_n(n, ith, nth); + int n_block_begin = n_start; + int n_block_size = n_end - n_block_begin; + if (n_block_size <= 0) return; + pack_block(src + (size_t)n_block_begin * src_stride, src_stride, n_block_begin, n_block_size); + } + void from_mat_transposed(ggml_bf16_t* src, int src_n, int src_k, int ith, int nth) { assert(n == src_k && k == src_n); auto [n_start, n_end] = K::split_range_n(n, ith, nth); diff --git a/kt-kernel/operators/amx/la/bf16_dweight.hpp b/kt-kernel/operators/amx/la/bf16_dweight.hpp new file mode 100644 index 000000000..02756fd9a --- /dev/null +++ b/kt-kernel/operators/amx/la/bf16_dweight.hpp @@ -0,0 +1,168 @@ +#ifndef AMX_BF16_DWEIGHT_HPP +#define AMX_BF16_DWEIGHT_HPP + +#include +#include +#include +#include +#include + +#include "amx_kernels.hpp" +#include "amx_raw_kernels.hpp" + +namespace amx { + +struct BF16DWeightTimings { + uint64_t pack_a_ns = 0; + uint64_t pack_a_calls = 0; + uint64_t pack_b_ns = 0; + uint64_t pack_b_calls = 0; + uint64_t kernel_gate_up_ns = 0; + uint64_t kernel_gate_up_calls = 0; + uint64_t kernel_down_ns = 0; + uint64_t kernel_down_calls = 0; + uint64_t store_ns = 0; + uint64_t store_calls = 0; + + void reset() { *this = {}; } +}; + +inline BF16DWeightTimings& bf16_dweight_timings() { + static thread_local BF16DWeightTimings timings; + return timings; +} + +class BF16DWeightScratch { + public: + using Kernel = GemmKernel224BF16; + static constexpr int M_STEP = Kernel::M_STEP; + static constexpr int N_STEP = Kernel::N_STEP; + + BF16DWeightScratch() = default; + BF16DWeightScratch(const BF16DWeightScratch&) = delete; + BF16DWeightScratch& operator=(const BF16DWeightScratch&) = delete; + + ~BF16DWeightScratch() { + std::free(a0_); + std::free(a1_); + std::free(b_); + } + + void ensure(int padded_k) { + if (padded_k <= capacity_k_) return; + const size_t a_elements = static_cast(M_STEP) * padded_k; + const size_t b_elements = static_cast(N_STEP) * padded_k; + resize(a0_, a_elements); + resize(a1_, a_elements); + resize(b_, b_elements); + capacity_k_ = padded_k; + } + + ggml_bf16_t* a0() { return a0_; } + ggml_bf16_t* a1() { return a1_; } + ggml_bf16_t* b() { return b_; } + float* c0() { return c0_; } + float* c1() { return c1_; } + + private: + static void resize(ggml_bf16_t*& buffer, size_t elements) { + void* replacement = nullptr; + if (posix_memalign(&replacement, 64, elements * sizeof(ggml_bf16_t)) != 0 || replacement == nullptr) { + throw std::runtime_error("failed to allocate BF16 dWeight scratch"); + } + std::free(buffer); + buffer = static_cast(replacement); + } + + int capacity_k_ = 0; + ggml_bf16_t* a0_ = nullptr; + ggml_bf16_t* a1_ = nullptr; + ggml_bf16_t* b_ = nullptr; + alignas(64) float c0_[M_STEP * N_STEP]; + alignas(64) float c1_[M_STEP * N_STEP]; +}; + +inline BF16DWeightScratch& bf16_dweight_scratch() { + static thread_local BF16DWeightScratch scratch; + return scratch; +} + +class BF16DWeightKernel { + public: + using Kernel = GemmKernel224BF16; + using BufferA = Kernel::BufferA; + using BufferB = Kernel::BufferB; + static constexpr int M_STEP = Kernel::M_STEP; + static constexpr int N_STEP = Kernel::N_STEP; + static constexpr int K_STEP = Kernel::K_STEP; + + static int padded_k(int routes) { return std::max(K_STEP, (routes + K_STEP - 1) / K_STEP * K_STEP); } + + static void configure_worker() { Kernel::config(); } + + static void pack_a_transposed(BufferA& destination, const ggml_bf16_t* source, int source_stride, int source_column, + int row_count, int routes) { + const int k = destination.k; + for (int k_begin = 0; k_begin < k; k_begin += K_STEP) { + ggml_bf16_t* tile = destination.get_submat(M_STEP, k, 0, k_begin); + std::memset(tile, 0, M_STEP * K_STEP * sizeof(ggml_bf16_t)); + const int valid_k = std::min(K_STEP, routes - k_begin); + if (valid_k <= 0) continue; + for (int row = 0; row < row_count; ++row) { + for (int kk = 0; kk < valid_k; ++kk) { + tile[row * K_STEP + kk] = source[static_cast(k_begin + kk) * source_stride + source_column + row]; + } + } + } + } + + static void pack_b_transposed(BufferB& destination, const ggml_bf16_t* source, int source_stride, int source_column, + int row_count, int routes) { + const int k = destination.k; + for (int k_begin = 0; k_begin < k; k_begin += K_STEP) { + ggml_bf16_t* tile = destination.get_submat(N_STEP, k, 0, k_begin); + std::memset(tile, 0, N_STEP * K_STEP * sizeof(ggml_bf16_t)); + const int valid_k = std::min(K_STEP, routes - k_begin); + if (valid_k > 0) { + for (int row = 0; row < row_count; ++row) { + for (int kk = 0; kk < valid_k; ++kk) { + tile[row * K_STEP + kk] = source[static_cast(k_begin + kk) * source_stride + source_column + row]; + } + } + } + transpose_16x16_32bit(reinterpret_cast<__m512i*>(tile)); + transpose_16x16_32bit(reinterpret_cast<__m512i*>(tile + Kernel::TILE_N * K_STEP)); + } + } + + static void multiply(int padded_k, float* destination, BufferA& a, BufferB& b) { + for (int k_block_begin = 0; k_block_begin < padded_k; k_block_begin += Kernel::K_BLOCK) { + if constexpr (AMX_AVAILABLE) { + Kernel::amx_kernel(M_STEP, N_STEP, padded_k, 0, 0, k_block_begin, destination, &a, &b); + } else { + Kernel::avx_kernel_4(M_STEP, N_STEP, padded_k, 0, 0, k_block_begin, destination, &a, &b); + } + } + } + + static void store_bf16(const float* source, ggml_bf16_t* destination, int destination_stride, int row_count, + int column_count) { + for (int row = 0; row < row_count; ++row) { + const float* src_row = source + row * N_STEP; + ggml_bf16_t* dst_row = destination + static_cast(row) * destination_stride; + if (column_count == N_STEP) { + __m512 lo = _mm512_loadu_ps(src_row); + __m512 hi = _mm512_loadu_ps(src_row + 16); + avx512_32xfp32_to_32xbf16(&lo, &hi, reinterpret_cast<__m512i*>(dst_row)); + } else { + for (int column = 0; column < column_count; ++column) { + dst_row[column] = GGML_FP32_TO_BF16(src_row[column]); + } + } + } + } +}; + +} // namespace amx + +#endif // AMX_BF16_DWEIGHT_HPP diff --git a/kt-kernel/operators/amx/sft_moe.hpp b/kt-kernel/operators/amx/sft_moe.hpp index 451732ad1..836f0ddac 100644 --- a/kt-kernel/operators/amx/sft_moe.hpp +++ b/kt-kernel/operators/amx/sft_moe.hpp @@ -29,6 +29,7 @@ #include "../../cpu_backend/worker_pool.h" #include "../sft_profile.hpp" #include "ggml.h" +#include "la/bf16_dweight.hpp" #include "la/amx_kernels.hpp" #include "la/amx_raw_kernels.hpp" #include "la/avx_kernels.hpp" @@ -362,6 +363,7 @@ template class BaseMOE = AMX_MOE_TP, bool SkipLoRA = class AMX_SFT_MOE_TP : public BaseMOE { public: static constexpr bool kSkipLoRA = SkipLoRA; + static constexpr bool kSupportsDirectBf16Reload = std::is_same_v; protected: using Base = BaseMOE; @@ -2004,15 +2006,135 @@ class AMX_SFT_MOE_TP : public BaseMOE { std::vector expert_offsets(activated_expert); size_t token_offset = 0; + int max_routes = 0; for (int task_id = 0; task_id < activated_expert; task_id++) { expert_offsets[task_id] = token_offset; int expert_idx = cache.m_expert_id_map_cache[task_id]; - token_offset += cache.m_local_num_cache[expert_idx]; + const int routes = cache.m_local_num_cache[expert_idx]; + token_offset += routes; + max_routes = std::max(max_routes, routes); } profiler_.record(SFTProfileStage::BwdBaseWeightGradOffsets, stage_start); - if constexpr ((std::is_same_v || std::is_same_v) && - amx::AMX_AVAILABLE) { + if constexpr (std::is_same_v) { + using DWeightKernel = amx::BF16DWeightKernel; + constexpr int TILE_M = DWeightKernel::M_STEP; + constexpr int TILE_N = DWeightKernel::N_STEP; + const int i_tiles = (I + TILE_M - 1) / TILE_M; + const int h_tiles = (H + TILE_N - 1) / TILE_N; + const int tasks_per_expert = i_tiles * 2; + const int total_tasks = activated_expert * tasks_per_expert; + const int max_padded_k = DWeightKernel::padded_k(max_routes); + auto pool = config_.pool->get_subpool(tp_part_idx); + const bool profile_inner = profiler_.enabled(); + + stage_start = profiler_.start(); + if (total_tasks > 0) { + pool->do_work_stealing_job( + total_tasks, + [max_padded_k](int _) { + DWeightKernel::configure_worker(); + amx::bf16_dweight_scratch().ensure(max_padded_k); + amx::bf16_dweight_timings().reset(); + }, + [&, i_tiles, h_tiles, tasks_per_expert](int task_id) { + const int expert_task = task_id / tasks_per_expert; + const int local_task = task_id % tasks_per_expert; + const bool do_down = local_task >= i_tiles; + const int i_tile = local_task % i_tiles; + const int expert_idx = cache.m_expert_id_map_cache[expert_task]; + const int routes = cache.m_local_num_cache[expert_idx]; + if (routes == 0) return; + + const size_t pos_start = expert_offsets[expert_task]; + const int padded_k = DWeightKernel::padded_k(routes); + const int i_start = i_tile * TILE_M; + const int i_count = std::min(TILE_M, I - i_start); + auto& scratch = amx::bf16_dweight_scratch(); + auto& timings = amx::bf16_dweight_timings(); + typename DWeightKernel::BufferA a0(TILE_M, padded_k, scratch.a0()); + typename DWeightKernel::BufferA a1(TILE_M, padded_k, scratch.a1()); + typename DWeightKernel::BufferB b(TILE_N, padded_k, scratch.b()); + + auto profile_operation = [profile_inner](uint64_t& elapsed_ns, uint64_t& calls, auto&& operation) { + if (!profile_inner) { + operation(); + return; + } + const auto begin = SFTProfiler::Clock::now(); + operation(); + elapsed_ns += static_cast( + std::chrono::duration_cast(SFTProfiler::Clock::now() - begin).count()); + calls++; + }; + + if (do_down) { + const ggml_bf16_t* grad_output = base_grad_output_bf16_ptr_[expert_idx]; + const ggml_bf16_t* intermediate = cache.intermediate_cache + pos_start * I; + profile_operation(timings.pack_b_ns, timings.pack_b_calls, [&] { + DWeightKernel::pack_b_transposed(b, intermediate, I, i_start, i_count, routes); + }); + + ggml_bf16_t* down_dst = gdp + static_cast(expert_idx) * H * F; + for (int h_tile = 0; h_tile < h_tiles; ++h_tile) { + const int h_start = h_tile * TILE_M; + const int h_count = std::min(TILE_M, H - h_start); + profile_operation(timings.pack_a_ns, timings.pack_a_calls, [&] { + DWeightKernel::pack_a_transposed(a0, grad_output, H, h_start, h_count, routes); + }); + profile_operation(timings.kernel_down_ns, timings.kernel_down_calls, + [&] { DWeightKernel::multiply(padded_k, scratch.c0(), a0, b); }); + profile_operation(timings.store_ns, timings.store_calls, [&] { + DWeightKernel::store_bf16(scratch.c0(), down_dst + static_cast(h_start) * F + i_start, F, + h_count, i_count); + }); + } + return; + } + + const ggml_bf16_t* gate_grad = grad_gate_output_ + pos_start * I; + const ggml_bf16_t* up_grad = grad_up_output_ + pos_start * I; + profile_operation(timings.pack_a_ns, timings.pack_a_calls, [&] { + DWeightKernel::pack_a_transposed(a0, gate_grad, I, i_start, i_count, routes); + DWeightKernel::pack_a_transposed(a1, up_grad, I, i_start, i_count, routes); + }); + + const ggml_bf16_t* input = m_local_input_ptr_[expert_idx]; + ggml_bf16_t* gate_dst = ggp + static_cast(expert_idx) * F * H; + ggml_bf16_t* up_dst = gup_ptr + static_cast(expert_idx) * F * H; + for (int h_tile = 0; h_tile < h_tiles; ++h_tile) { + const int h_start = h_tile * TILE_N; + const int h_count = std::min(TILE_N, H - h_start); + profile_operation(timings.pack_b_ns, timings.pack_b_calls, + [&] { DWeightKernel::pack_b_transposed(b, input, H, h_start, h_count, routes); }); + profile_operation(timings.kernel_gate_up_ns, timings.kernel_gate_up_calls, [&] { + DWeightKernel::multiply(padded_k, scratch.c0(), a0, b); + DWeightKernel::multiply(padded_k, scratch.c1(), a1, b); + }); + profile_operation(timings.store_ns, timings.store_calls, [&] { + DWeightKernel::store_bf16(scratch.c0(), gate_dst + static_cast(i_start) * H + h_start, H, + i_count, h_count); + DWeightKernel::store_bf16(scratch.c1(), up_dst + static_cast(i_start) * H + h_start, H, + i_count, h_count); + }); + } + }, + [this](int _) { + auto& timings = amx::bf16_dweight_timings(); + profiler_.record_ns(SFTProfileStage::BwdBaseWeightGradPackA, timings.pack_a_ns, timings.pack_a_calls); + profiler_.record_ns(SFTProfileStage::BwdBaseWeightGradPackB, timings.pack_b_ns, timings.pack_b_calls); + profiler_.record_ns(SFTProfileStage::BwdBaseWeightGradKernelGateUp, timings.kernel_gate_up_ns, + timings.kernel_gate_up_calls); + profiler_.record_ns(SFTProfileStage::BwdBaseWeightGradKernelDown, timings.kernel_down_ns, + timings.kernel_down_calls); + profiler_.record_ns(SFTProfileStage::BwdBaseWeightGradStore, timings.store_ns, timings.store_calls); + }); + } + profiler_.record(SFTProfileStage::BwdBaseWeightGradMatMat, stage_start); + return; + } + + if constexpr (std::is_same_v && amx::AMX_AVAILABLE) { constexpr int TILE_M = T::M_STEP; constexpr int TILE_N = T::N_STEP; constexpr int TILE_K = T::K_STEP; @@ -2677,6 +2799,54 @@ class AMX_SFT_MOE_TP : public BaseMOE { backward_weights_prepared_ = true; } + void load_forward_weights_from_full_bf16(void* gate_proj, void* up_proj, void* down_proj, int full_intermediate_size, + int intermediate_offset) { + if constexpr (!kSupportsDirectBf16Reload) { + throw std::runtime_error("direct BF16 reload requires GemmKernel224BF16"); + } else { + const int H = config_.hidden_size; + const int I = config_.intermediate_size; + if (gate_proj == nullptr || up_proj == nullptr || down_proj == nullptr) { + throw std::runtime_error("direct BF16 reload requires all three base weights"); + } + if (intermediate_offset < 0 || full_intermediate_size < intermediate_offset + I) { + throw std::runtime_error("invalid TP slice for direct BF16 reload"); + } + + const auto* physical_to_logical_map = static_cast(config_.physical_to_logical_map); + auto* gate = static_cast(gate_proj); + auto* up = static_cast(up_proj); + auto* down = static_cast(down_proj); + auto pool = config_.pool->get_subpool(tp_part_idx); + + const int gate_up_nth = T::recommended_nth(I); + pool->do_work_stealing_job( + gate_up_nth * config_.expert_num, nullptr, + [&, gate_up_nth, physical_to_logical_map](int task_id) { + const int physical_expert = task_id / gate_up_nth; + const int ith = task_id % gate_up_nth; + const size_t logical_expert = expert_map(physical_to_logical_map, physical_expert); + const size_t source_offset = + logical_expert * full_intermediate_size * H + static_cast(intermediate_offset) * H; + gate_bb_[physical_expert]->from_mat_strided(gate + source_offset, H, ith, gate_up_nth); + up_bb_[physical_expert]->from_mat_strided(up + source_offset, H, ith, gate_up_nth); + }, + nullptr); + + const int down_nth = T::recommended_nth(H); + pool->do_work_stealing_job( + down_nth * config_.expert_num, nullptr, + [&, down_nth, physical_to_logical_map](int task_id) { + const int physical_expert = task_id / down_nth; + const int ith = task_id % down_nth; + const size_t logical_expert = expert_map(physical_to_logical_map, physical_expert); + const size_t source_offset = logical_expert * H * full_intermediate_size + intermediate_offset; + down_bb_[physical_expert]->from_mat_strided(down + source_offset, full_intermediate_size, ith, down_nth); + }, + nullptr); + } + } + /** * @brief Standalone method for async backward BB repack (Phase 2). * Called from TP_MOE_SFT::submit_backward_repack() on a separate thread. diff --git a/kt-kernel/operators/amx/test/test_bf16_dweight.cpp b/kt-kernel/operators/amx/test/test_bf16_dweight.cpp new file mode 100644 index 000000000..fd90499fb --- /dev/null +++ b/kt-kernel/operators/amx/test/test_bf16_dweight.cpp @@ -0,0 +1,106 @@ +#include +#include +#include +#include +#include +#include + +#include "../la/bf16_dweight.hpp" + +namespace { + +using DWeightKernel = amx::BF16DWeightKernel; +using Kernel = DWeightKernel::Kernel; + +void* alloc_buffer(size_t bytes) { + void* pointer = nullptr; + if (posix_memalign(&pointer, 64, (bytes + 63) / 64 * 64) != 0 || pointer == nullptr) std::abort(); + std::memset(pointer, 0, bytes); + return pointer; +} + +void fill_random(std::vector& values, unsigned seed) { + std::mt19937 generator(seed); + std::uniform_real_distribution distribution(-0.25f, 0.25f); + for (auto& value : values) value = GGML_FP32_TO_BF16(distribution(generator)); +} + +bool run_case(int routes, int rows, int columns) { + constexpr int source_column = 2; + constexpr int destination_column = 3; + const int lhs_stride = rows + source_column + 3; + const int rhs_stride = columns + source_column + 5; + const int destination_stride = columns + destination_column + 7; + const int padded_k = DWeightKernel::padded_k(routes); + + std::vector lhs(static_cast(routes) * lhs_stride); + std::vector rhs(static_cast(routes) * rhs_stride); + std::vector actual(static_cast(rows) * destination_stride); + std::vector expected(static_cast(rows) * columns); + fill_random(lhs, static_cast(routes * 17 + rows)); + fill_random(rhs, static_cast(routes * 31 + columns)); + + void* a_memory = alloc_buffer(Kernel::BufferA::required_size(Kernel::M_STEP, padded_k)); + void* b_memory = alloc_buffer(Kernel::BufferB::required_size(Kernel::N_STEP, padded_k)); + Kernel::BufferA a(Kernel::M_STEP, padded_k, a_memory); + Kernel::BufferB b(Kernel::N_STEP, padded_k, b_memory); + alignas(64) float accumulator[Kernel::M_STEP * Kernel::N_STEP]; + + DWeightKernel::pack_a_transposed(a, lhs.data(), lhs_stride, source_column, rows, routes); + DWeightKernel::pack_b_transposed(b, rhs.data(), rhs_stride, source_column, columns, routes); + DWeightKernel::multiply(padded_k, accumulator, a, b); + DWeightKernel::store_bf16(accumulator, actual.data() + destination_column, destination_stride, rows, columns); + + for (int row = 0; row < rows; ++row) { + for (int column = 0; column < columns; ++column) { + float sum = 0.0f; + for (int route = 0; route < routes; ++route) { + sum += GGML_BF16_TO_FP32(lhs[static_cast(route) * lhs_stride + source_column + row]) * + GGML_BF16_TO_FP32(rhs[static_cast(route) * rhs_stride + source_column + column]); + } + expected[static_cast(row) * columns + column] = GGML_FP32_TO_BF16(sum); + } + } + + double difference_sq = 0.0; + double expected_sq = 0.0; + double actual_sq = 0.0; + double dot = 0.0; + float max_abs = 0.0f; + for (int row = 0; row < rows; ++row) { + for (int column = 0; column < columns; ++column) { + const float expected_value = GGML_BF16_TO_FP32(expected[static_cast(row) * columns + column]); + const float actual_value = + GGML_BF16_TO_FP32(actual[static_cast(row) * destination_stride + destination_column + column]); + const double difference = static_cast(actual_value) - expected_value; + difference_sq += difference * difference; + expected_sq += static_cast(expected_value) * expected_value; + actual_sq += static_cast(actual_value) * actual_value; + dot += static_cast(expected_value) * actual_value; + max_abs = std::max(max_abs, std::fabs(actual_value - expected_value)); + } + } + + const double relative_l2 = std::sqrt(difference_sq / std::max(expected_sq, 1e-30)); + const double cosine = dot / std::sqrt(std::max(expected_sq * actual_sq, 1e-30)); + const bool passed = relative_l2 <= 0.01 && cosine >= 0.999; + std::printf("BF16 dWeight routes=%d shape=%dx%d: rel_l2=%.6e cosine=%.9f max_abs=%.6e %s\n", routes, rows, columns, + relative_l2, cosine, max_abs, passed ? "PASS" : "FAIL"); + + std::free(a_memory); + std::free(b_memory); + return passed; +} + +} // namespace + +int main() { + DWeightKernel::configure_worker(); + bool passed = true; + for (int routes : {1, 31, 32, 33, 65, 1792, 1825}) { + passed = run_case(routes, 32, 32) && passed; + } + passed = run_case(33, 17, 29) && passed; + passed = run_case(65, 31, 7) && passed; + return passed ? 0 : 1; +} diff --git a/kt-kernel/operators/amx/test/test_raw_bf16_repack.cpp b/kt-kernel/operators/amx/test/test_raw_bf16_repack.cpp index 50bf9853d..55a13097e 100644 --- a/kt-kernel/operators/amx/test/test_raw_bf16_repack.cpp +++ b/kt-kernel/operators/amx/test/test_raw_bf16_repack.cpp @@ -30,6 +30,11 @@ void from_mat(BufferB& buffer, ggml_bf16_t* source) { for (int ith = 0; ith < nth; ith++) buffer.from_mat(source, ith, nth); } +void from_mat_strided(BufferB& buffer, ggml_bf16_t* source, int source_stride) { + int nth = Kernel::recommended_nth(buffer.n); + for (int ith = 0; ith < nth; ith++) buffer.from_mat_strided(source, source_stride, ith, nth); +} + void to_mat(const BufferB& buffer, ggml_bf16_t* destination) { int nth = Kernel::recommended_nth(buffer.n); for (int ith = 0; ith < nth; ith++) buffer.to_mat(destination, ith, nth); @@ -87,6 +92,34 @@ bool run_case(int n, int k) { return roundtrip_ok && transpose_ok && direct_ok; } +bool run_strided_case(int n, int k, int source_stride, size_t source_offset) { + const size_t source_count = source_offset + (size_t)(n - 1) * source_stride + k; + const size_t output_count = (size_t)n * k; + std::vector source(source_count); + std::vector output(output_count); + fill_random(source, (unsigned)(n * 17 + k * 13 + source_stride)); + + void* memory = alloc_buffer(BufferB::required_size(n, k)); + BufferB packed(n, k, memory); + from_mat_strided(packed, source.data() + source_offset, source_stride); + to_mat(packed, output.data()); + + bool passed = true; + for (int row = 0; row < n && passed; row++) { + for (int column = 0; column < k; column++) { + if (source[source_offset + (size_t)row * source_stride + column].bits != output[(size_t)row * k + column].bits) { + passed = false; + break; + } + } + } + + std::free(memory); + std::printf("raw BF16 strided repack %dx%d stride=%d offset=%zu: %s\n", n, k, source_stride, source_offset, + passed ? "PASS" : "FAIL"); + return passed; +} + } // namespace int main() { @@ -94,5 +127,7 @@ int main() { passed = run_case(64, 64) && passed; passed = run_case(768, 2048) && passed; passed = run_case(2048, 768) && passed; + passed = run_strided_case(64, 64, 64, 32 * 64) && passed; + passed = run_strided_case(64, 64, 96, 17) && passed; return passed ? 0 : 1; } diff --git a/kt-kernel/operators/moe-sft-tp.hpp b/kt-kernel/operators/moe-sft-tp.hpp index 4605d6c7b..12ca17cb8 100644 --- a/kt-kernel/operators/moe-sft-tp.hpp +++ b/kt-kernel/operators/moe-sft-tp.hpp @@ -369,80 +369,112 @@ class TP_MOE_SFT : public TP_MOE { throw std::runtime_error("K2 pre-quantized mode does not support TP > 1 yet"); } } else if (config.gate_proj != nullptr) { - // printf("TP_MOE_SFT: From BF16 with partitioning\n"); - auto reload_stage_start = profiler_.start(); + if constexpr (T::kSupportsDirectBf16Reload) { + std::vector intermediate_offsets(tp_count); + int intermediate_offset = 0; + for (int i = 0; i < tp_count; ++i) { + const auto& tpc = tp_configs[i]; + if (tpc.hidden_size != config.hidden_size || tpc.expert_num != config.expert_num) { + throw std::runtime_error("incompatible TP config for direct BF16 reload"); + } + intermediate_offsets[i] = intermediate_offset; + intermediate_offset += tpc.intermediate_size; + tps[i]->set_physical_to_logical_map(config.physical_to_logical_map); + } + if (intermediate_offset != config.intermediate_size) { + throw std::runtime_error("TP intermediate slices do not cover the full BF16 weight"); + } - // Temporary storage for partitioned weights - std::vector temp_gate(tp_count); - std::vector temp_up(tp_count); - std::vector temp_down(tp_count); + auto reload_stage_start = profiler_.start(); + pool->dispense_backend()->do_numa_job([&, this](int numa_id) { + tps[numa_id]->load_forward_weights_from_full_bf16(config.gate_proj, config.up_proj, config.down_proj, + config.intermediate_size, intermediate_offsets[numa_id]); + }); + profiler_.record(SFTProfileStage::BaseWeightReloadDirectPack, reload_stage_start); - // Step 1: For each NUMA, allocate and copy partitioned weights - for (int i = 0; i < tp_count; i++) { - // Use tp_configs[i] instead of tps[i]->config_ (which is protected) - auto& tpc = tp_configs[i]; - size_t gate_up_elcount = (size_t)tpc.intermediate_size * tpc.hidden_size; - - // Allocate partitioned weight space - temp_gate[i] = new ggml_bf16_t[tpc.expert_num * gate_up_elcount]; - temp_up[i] = new ggml_bf16_t[tpc.expert_num * gate_up_elcount]; - temp_down[i] = new ggml_bf16_t[tpc.expert_num * gate_up_elcount]; - - // Copy partitioned weights - pool->get_subpool(i)->do_work_stealing_job( - tpc.expert_num, nullptr, - [&, i, gate_up_elcount](int expert_id_) { - size_t expert_id = expert_map(physical_to_logical_map, expert_id_); - - // gate_proj/up_proj: [intermediate_size, hidden_size] - contiguous block slice - memcpy(temp_gate[i] + expert_id * gate_up_elcount, - (ggml_bf16_t*)config.gate_proj + expert_id * config.intermediate_size * config.hidden_size + - i * gate_up_elcount, - sizeof(ggml_bf16_t) * gate_up_elcount); - - memcpy(temp_up[i] + expert_id * gate_up_elcount, - (ggml_bf16_t*)config.up_proj + expert_id * config.intermediate_size * config.hidden_size + - i * gate_up_elcount, - sizeof(ggml_bf16_t) * gate_up_elcount); - - // down_proj: [hidden_size, intermediate_size] - row-wise slice - for (size_t col = 0; col < config.hidden_size; col++) { - memcpy(temp_down[i] + expert_id * tpc.hidden_size * tpc.intermediate_size + col * tpc.intermediate_size, - (ggml_bf16_t*)config.down_proj + expert_id * config.intermediate_size * config.hidden_size + - col * config.intermediate_size + i * tpc.intermediate_size, - sizeof(ggml_bf16_t) * tpc.intermediate_size); - } - }, - nullptr); - } - profiler_.record(SFTProfileStage::BaseWeightReloadPartition, reload_stage_start); + if (!config.share_backward_bb) { + reload_stage_start = profiler_.start(); + pool->dispense_backend()->do_numa_job( + [this](int numa_id) { tps[numa_id]->prepare_backward_weights_from_forward(); }); + profiler_.record(SFTProfileStage::BaseWeightReloadBackwardPack, reload_stage_start); + } + } else { + // printf("TP_MOE_SFT: From BF16 with partitioning\n"); + auto reload_stage_start = profiler_.start(); - // Step 2: Set weight pointers BEFORE load_weights (Bug #24 fix) - reload_stage_start = profiler_.start(); - for (int i = 0; i < tp_count; i++) { - tps[i]->set_physical_to_logical_map(config.physical_to_logical_map); - tps[i]->set_weight_pointers_for_forward(temp_gate[i], temp_up[i], temp_down[i]); - } + // Temporary storage for partitioned weights + std::vector temp_gate(tp_count); + std::vector temp_up(tp_count); + std::vector temp_down(tp_count); - pool->dispense_backend()->do_numa_job([this](int numa_id) { tps[numa_id]->load_weights(); }); - profiler_.record(SFTProfileStage::BaseWeightReloadForwardPack, reload_stage_start); + // Step 1: For each NUMA, allocate and copy partitioned weights + for (int i = 0; i < tp_count; i++) { + // Use tp_configs[i] instead of tps[i]->config_ (which is protected) + auto& tpc = tp_configs[i]; + size_t gate_up_elcount = (size_t)tpc.intermediate_size * tpc.hidden_size; - // Step 3: Prepare backward weights (this also clears weight pointers) - reload_stage_start = profiler_.start(); - for (int i = 0; i < tp_count; i++) { - if (!config.share_backward_bb) { - tps[i]->prepare_bwd(temp_gate[i], temp_up[i], temp_down[i]); + // Allocate partitioned weight space + temp_gate[i] = new ggml_bf16_t[tpc.expert_num * gate_up_elcount]; + temp_up[i] = new ggml_bf16_t[tpc.expert_num * gate_up_elcount]; + temp_down[i] = new ggml_bf16_t[tpc.expert_num * gate_up_elcount]; + + // Copy partitioned weights + pool->get_subpool(i)->do_work_stealing_job( + tpc.expert_num, nullptr, + [&, i, gate_up_elcount](int expert_id_) { + size_t expert_id = expert_map(physical_to_logical_map, expert_id_); + + // gate_proj/up_proj: [intermediate_size, hidden_size] - contiguous block slice + memcpy(temp_gate[i] + expert_id * gate_up_elcount, + (ggml_bf16_t*)config.gate_proj + expert_id * config.intermediate_size * config.hidden_size + + i * gate_up_elcount, + sizeof(ggml_bf16_t) * gate_up_elcount); + + memcpy(temp_up[i] + expert_id * gate_up_elcount, + (ggml_bf16_t*)config.up_proj + expert_id * config.intermediate_size * config.hidden_size + + i * gate_up_elcount, + sizeof(ggml_bf16_t) * gate_up_elcount); + + // down_proj: [hidden_size, intermediate_size] - row-wise slice + for (size_t col = 0; col < config.hidden_size; col++) { + memcpy( + temp_down[i] + expert_id * tpc.hidden_size * tpc.intermediate_size + col * tpc.intermediate_size, + (ggml_bf16_t*)config.down_proj + expert_id * config.intermediate_size * config.hidden_size + + col * config.intermediate_size + i * tpc.intermediate_size, + sizeof(ggml_bf16_t) * tpc.intermediate_size); + } + }, + nullptr); } - } - profiler_.record(SFTProfileStage::BaseWeightReloadBackwardPack, reload_stage_start); + profiler_.record(SFTProfileStage::BaseWeightReloadPartition, reload_stage_start); - reload_stage_start = profiler_.start(); - for (int i = 0; i < tp_count; i++) { - delete[] (temp_gate[i]); - delete[] (temp_up[i]); - delete[] (temp_down[i]); + // Step 2: Set weight pointers BEFORE load_weights (Bug #24 fix) + reload_stage_start = profiler_.start(); + for (int i = 0; i < tp_count; i++) { + tps[i]->set_physical_to_logical_map(config.physical_to_logical_map); + tps[i]->set_weight_pointers_for_forward(temp_gate[i], temp_up[i], temp_down[i]); + } + + pool->dispense_backend()->do_numa_job([this](int numa_id) { tps[numa_id]->load_weights(); }); + profiler_.record(SFTProfileStage::BaseWeightReloadForwardPack, reload_stage_start); + + // Step 3: Prepare backward weights (this also clears weight pointers) + reload_stage_start = profiler_.start(); + for (int i = 0; i < tp_count; i++) { + if (!config.share_backward_bb) { + tps[i]->prepare_bwd(temp_gate[i], temp_up[i], temp_down[i]); + } + } + profiler_.record(SFTProfileStage::BaseWeightReloadBackwardPack, reload_stage_start); + + reload_stage_start = profiler_.start(); + for (int i = 0; i < tp_count; i++) { + delete[] (temp_gate[i]); + delete[] (temp_up[i]); + delete[] (temp_down[i]); + } + profiler_.record(SFTProfileStage::BaseWeightReloadCleanup, reload_stage_start); } - profiler_.record(SFTProfileStage::BaseWeightReloadCleanup, reload_stage_start); } else { // Other loading methods (from loader or file) for (int i = 0; i < tp_count; i++) { diff --git a/kt-kernel/operators/sft_profile.hpp b/kt-kernel/operators/sft_profile.hpp index 91ddefb77..e760c9d9b 100644 --- a/kt-kernel/operators/sft_profile.hpp +++ b/kt-kernel/operators/sft_profile.hpp @@ -51,6 +51,11 @@ enum class SFTProfileStage : uint8_t { BwdRouterGrad, BwdBaseWeightGrad, BwdBaseWeightGradOffsets, + BwdBaseWeightGradMatMat, + BwdBaseWeightGradPackA, + BwdBaseWeightGradPackB, + BwdBaseWeightGradKernelGateUp, + BwdBaseWeightGradKernelDown, BwdBaseWeightGradAmx, BwdBaseWeightGradZero, BwdBaseWeightGradGateUp, @@ -71,6 +76,7 @@ enum class SFTProfileStage : uint8_t { BaseWeightReload, BaseWeightReloadPartition, BaseWeightReloadForwardPack, + BaseWeightReloadDirectPack, BaseWeightReloadBackwardPack, BaseWeightReloadCleanup, @@ -112,6 +118,11 @@ inline constexpr std::array(SFTProfileStage::Co "backward.router_grad", "backward.base_weight_grad", "backward.base_weight_grad.offsets", + "backward.base_weight_grad.matmat", + "backward.base_weight_grad.worker_cpu.pack_a", + "backward.base_weight_grad.worker_cpu.pack_b", + "backward.base_weight_grad.worker_cpu.kernel_gate_up", + "backward.base_weight_grad.worker_cpu.kernel_down", "backward.base_weight_grad.amx", "backward.base_weight_grad.zero", "backward.base_weight_grad.gate_up", @@ -130,14 +141,15 @@ inline constexpr std::array(SFTProfileStage::Co "weights.base_reload", "weights.base_reload.partition", "weights.base_reload.forward_pack", + "weights.base_reload.direct_pack", "weights.base_reload.backward_pack", "weights.base_reload.cleanup", }; inline bool sft_profile_enabled_from_env() { const char* value = std::getenv("KT_SFT_PROFILE"); - return value != nullptr && value[0] != '\0' && std::strcmp(value, "0") != 0 && - std::strcmp(value, "false") != 0 && std::strcmp(value, "False") != 0; + return value != nullptr && value[0] != '\0' && std::strcmp(value, "0") != 0 && std::strcmp(value, "false") != 0 && + std::strcmp(value, "False") != 0; } class SFTProfiler { @@ -145,9 +157,7 @@ class SFTProfiler { using Clock = std::chrono::steady_clock; using TimePoint = Clock::time_point; - explicit SFTProfiler(bool enabled = sft_profile_enabled_from_env()) : enabled_(enabled) { - reset(); - } + explicit SFTProfiler(bool enabled = sft_profile_enabled_from_env()) : enabled_(enabled) { reset(); } bool enabled() const { return enabled_; } @@ -156,9 +166,14 @@ class SFTProfiler { void record(SFTProfileStage stage, TimePoint start) { if (!enabled_) return; const auto elapsed = std::chrono::duration_cast(Clock::now() - start).count(); + record_ns(stage, static_cast(elapsed)); + } + + void record_ns(SFTProfileStage stage, uint64_t elapsed_ns, uint64_t calls = 1) { + if (!enabled_) return; const size_t idx = static_cast(stage); - total_ns_[idx].fetch_add(static_cast(elapsed), std::memory_order_relaxed); - calls_[idx].fetch_add(1, std::memory_order_relaxed); + total_ns_[idx].fetch_add(elapsed_ns, std::memory_order_relaxed); + calls_[idx].fetch_add(calls, std::memory_order_relaxed); } void record_workload(uint64_t tokens, uint64_t routed_rows, uint64_t active_experts) { diff --git a/kt-kernel/python/sft/profiler.py b/kt-kernel/python/sft/profiler.py index 6d1bb8336..e909166cd 100644 --- a/kt-kernel/python/sft/profiler.py +++ b/kt-kernel/python/sft/profiler.py @@ -63,6 +63,8 @@ def _split_timer_key(key: str) -> tuple[str, str] | None: def _parent_stage(stage: str) -> str | None: + if stage.startswith("backward.base_weight_grad.worker_cpu."): + return None if stage.startswith("backward.base_weight_grad."): return "backward.base_weight_grad" if stage.startswith("weights.base_reload."): From ea84e6edf44d2576dc3078d60f5b2b808b42d095 Mon Sep 17 00:00:00 2001 From: yyj Date: Fri, 17 Jul 2026 15:03:54 +0800 Subject: [PATCH 14/20] [test](kt-kernel): benchmark BF16 dWeight AMX driver --- .../operators/amx/test/test_bf16_dweight.cpp | 76 ++++++++++++++++++- 1 file changed, 75 insertions(+), 1 deletion(-) diff --git a/kt-kernel/operators/amx/test/test_bf16_dweight.cpp b/kt-kernel/operators/amx/test/test_bf16_dweight.cpp index fd90499fb..2e845aade 100644 --- a/kt-kernel/operators/amx/test/test_bf16_dweight.cpp +++ b/kt-kernel/operators/amx/test/test_bf16_dweight.cpp @@ -1,3 +1,5 @@ +#include +#include #include #include #include @@ -92,10 +94,82 @@ bool run_case(int routes, int rows, int columns) { return passed; } +bool run_amx_benchmark() { + if constexpr (!amx::AMX_AVAILABLE) { + std::printf("BF16 dWeight AMX benchmark: SKIP (AMX unavailable)\n"); + return true; + } + + constexpr int routes = 1024; + constexpr int iterations = 20000; + constexpr int rounds = 7; + const int padded_k = DWeightKernel::padded_k(routes); + std::vector lhs(static_cast(routes) * Kernel::M_STEP); + std::vector rhs(static_cast(routes) * Kernel::N_STEP); + fill_random(lhs, 20260717); + fill_random(rhs, 20260718); + + void* a_memory = alloc_buffer(Kernel::BufferA::required_size(Kernel::M_STEP, padded_k)); + void* b_memory = alloc_buffer(Kernel::BufferB::required_size(Kernel::N_STEP, padded_k)); + Kernel::BufferA a(Kernel::M_STEP, padded_k, a_memory); + Kernel::BufferB b(Kernel::N_STEP, padded_k, b_memory); + alignas(64) float accumulator[Kernel::M_STEP * Kernel::N_STEP]; + DWeightKernel::pack_a_transposed(a, lhs.data(), Kernel::M_STEP, 0, Kernel::M_STEP, routes); + DWeightKernel::pack_b_transposed(b, rhs.data(), Kernel::N_STEP, 0, Kernel::N_STEP, routes); + + auto legacy_tile_loop = [&] { + Kernel::clean_c(); + for (int k_begin = 0; k_begin < padded_k; k_begin += Kernel::K_STEP) { + Kernel::load_b(b.get_submat(Kernel::N_STEP, padded_k, 0, k_begin), + Kernel::K_STEP * sizeof(ggml_bf16_t)); + Kernel::load_a(a.get_submat(Kernel::M_STEP, padded_k, 0, k_begin), + Kernel::K_STEP * sizeof(ggml_bf16_t)); + Kernel::run_tile(); + } + Kernel::store_c(accumulator, Kernel::N_STEP * sizeof(float)); + }; + auto common_driver = [&] { DWeightKernel::multiply(padded_k, accumulator, a, b); }; + + for (int warmup = 0; warmup < 200; ++warmup) { + legacy_tile_loop(); + common_driver(); + } + + auto measure = [&](auto&& operation) { + const auto begin = std::chrono::steady_clock::now(); + for (int iteration = 0; iteration < iterations; ++iteration) operation(); + return std::chrono::duration(std::chrono::steady_clock::now() - begin).count() / iterations; + }; + std::vector legacy_ns; + std::vector common_ns; + for (int round = 0; round < rounds; ++round) { + if (round % 2 == 0) { + legacy_ns.push_back(measure(legacy_tile_loop)); + common_ns.push_back(measure(common_driver)); + } else { + common_ns.push_back(measure(common_driver)); + legacy_ns.push_back(measure(legacy_tile_loop)); + } + } + std::sort(legacy_ns.begin(), legacy_ns.end()); + std::sort(common_ns.begin(), common_ns.end()); + const double legacy_median = legacy_ns[rounds / 2]; + const double common_median = common_ns[rounds / 2]; + const double ratio = common_median / legacy_median; + const bool passed = ratio <= 1.05; + std::printf("BF16 dWeight AMX kernel routes=%d: legacy=%.1f ns common=%.1f ns ratio=%.4f %s\n", routes, + legacy_median, common_median, ratio, passed ? "PASS" : "FAIL"); + + std::free(a_memory); + std::free(b_memory); + return passed; +} + } // namespace -int main() { +int main(int argc, char** argv) { DWeightKernel::configure_worker(); + if (argc == 2 && std::strcmp(argv[1], "--benchmark") == 0) return run_amx_benchmark() ? 0 : 1; bool passed = true; for (int routes : {1, 31, 32, 33, 65, 1792, 1825}) { passed = run_case(routes, 32, 32) && passed; From 273e670890c7b966ce537cd1bb4acf69c8202767 Mon Sep 17 00:00:00 2001 From: yyj Date: Fri, 17 Jul 2026 15:17:41 +0800 Subject: [PATCH 15/20] [fix](kt-kernel): label dWeight store as worker CPU time --- kt-kernel/operators/sft_profile.hpp | 2 +- kt-kernel/test/per_commit/test_sft_profiler.py | 6 ++++++ 2 files changed, 7 insertions(+), 1 deletion(-) diff --git a/kt-kernel/operators/sft_profile.hpp b/kt-kernel/operators/sft_profile.hpp index e760c9d9b..07548e1bd 100644 --- a/kt-kernel/operators/sft_profile.hpp +++ b/kt-kernel/operators/sft_profile.hpp @@ -127,7 +127,7 @@ inline constexpr std::array(SFTProfileStage::Co "backward.base_weight_grad.zero", "backward.base_weight_grad.gate_up", "backward.base_weight_grad.down", - "backward.base_weight_grad.store", + "backward.base_weight_grad.worker_cpu.store", "tp.forward.total", "tp.forward.numa_compute", "tp.forward.merge", diff --git a/kt-kernel/test/per_commit/test_sft_profiler.py b/kt-kernel/test/per_commit/test_sft_profiler.py index fbe5935b8..f411f59df 100644 --- a/kt-kernel/test/per_commit/test_sft_profiler.py +++ b/kt-kernel/test/per_commit/test_sft_profiler.py @@ -29,6 +29,10 @@ def get_profile_stats(self, reset=False): "tp.0.backward.total.calls": 1, "tp.0.backward.down.total.total_ns": 2_000_000, "tp.0.backward.down.total.calls": 1, + "tp.0.backward.base_weight_grad.total_ns": 1_000_000, + "tp.0.backward.base_weight_grad.calls": 1, + "tp.0.backward.base_weight_grad.worker_cpu.store.total_ns": 2_000_000, + "tp.0.backward.base_weight_grad.worker_cpu.store.calls": 4, } def reset_profile_stats(self): @@ -55,6 +59,8 @@ def test_collect_and_format_profile(): assert "75.0%" in output assert "10.0%" in output assert "40.0%" in output + worker_store = next(line for line in output.splitlines() if "worker_cpu.store" in line) + assert worker_store.endswith("0.0%") def test_reset_and_disabled_profile(): From 61e63a8dd242b21d62649da8a4cf999033c095eb Mon Sep 17 00:00:00 2001 From: yyj Date: Fri, 17 Jul 2026 15:17:41 +0800 Subject: [PATCH 16/20] [docs](sft): record BF16 Full-FT performance --- doc/SUMMARY.md | 1 + .../Qwen3-30B-A3B-Full-FT-BF16-Performance.md | 113 ++++++++++++++++++ doc/en/SFT/README.md | 1 + 3 files changed, 115 insertions(+) create mode 100644 doc/en/SFT/Qwen3-30B-A3B-Full-FT-BF16-Performance.md diff --git a/doc/SUMMARY.md b/doc/SUMMARY.md index d19f29ca3..7944d71d5 100644 --- a/doc/SUMMARY.md +++ b/doc/SUMMARY.md @@ -11,6 +11,7 @@ - [KT-FT 微调推理闭环](zh/Qwen3.5-SGLang-LoRA-Serving_zh.md) - [Injection Tutorial](en/SFT/injection_tutorial.md) - [kt-sft developer tech notes](en/SFT/KTransformers-Fine-Tuning_Developer-Technical-Notes.md) + - [Qwen3-30B-A3B BF16 Full-FT Performance](en/SFT/Qwen3-30B-A3B-Full-FT-BF16-Performance.md) - [DPO tutorial](en/SFT/DPO_tutorial.md) diff --git a/doc/en/SFT/Qwen3-30B-A3B-Full-FT-BF16-Performance.md b/doc/en/SFT/Qwen3-30B-A3B-Full-FT-BF16-Performance.md new file mode 100644 index 000000000..8b1d68c54 --- /dev/null +++ b/doc/en/SFT/Qwen3-30B-A3B-Full-FT-BF16-Performance.md @@ -0,0 +1,113 @@ +# Qwen3-30B-A3B BF16 Full-FT Performance + +## Summary + +This work removes two CPU bottlenecks in raw-BF16 routed-expert full-parameter training: + +1. Expert weight gradients now aggregate all routed rows into BF16 mat-mat operations. The common driver dispatches to AVX512-BF16 on non-AMX CPUs and to the existing AMX microkernel on AMX CPUs. +2. Updated full BF16 expert weights are packed directly into persistent TP-local forward buffers. The reload no longer allocates, copies, and frees six temporary TP partitions per layer and step. + +On a dual-socket AMD EPYC 9355 with one RTX 5090, Qwen3-30B-A3B at batch 1 and sequence length 1024 improves from 22.653 to 200.701 token/s. The comparison uses steps 8-15 of two 15-step runs. + +## Implementation + +The main implementation is commit `34d2102` on `fullft-development`: + +- `BF16DWeightKernel` packs route-major activation and gradient panels into the existing raw BF16 `BufferA` and `BufferB` layouts. +- Work is split by `expert x intermediate tile x {gate/up, down}` and scheduled through the NUMA-local work-stealing pool. +- Gate and up share the packed input panel. Down reuses its packed intermediate panel. +- FP32 accumulation and BF16 gradient storage are unchanged. +- `BufferB::from_mat_strided()` packs a TP slice directly from the full `[E,F,H]` or `[E,H,F]` parameter tensor. +- The old scalar fallback remains available for non-BF16 kernels. + +The staged profiler distinguishes wall-clock stages from summed worker CPU time. Worker stages use the `backward.base_weight_grad.worker_cpu.*` namespace and intentionally have no wall-clock parent percentage. + +## Correctness + +| Check | Result | +|---|---| +| AVX512-BF16 dW routes `1,31,32,33,65,1792,1825` and tail shapes | Exact match after BF16 store | +| Contiguous, transposed, and strided raw-BF16 repack | Pass | +| Synthetic full-FT gradients, TP1 and TP2, qlen `8,31,32,33,65` | 30/30 pass | +| Real HF layer-0 expert-0 gradients, TP1 and TP2 | Relative L2 `0.00355-0.00378`, cosine approximately 1 | +| 15-step end-to-end training | Finite loss and grad norm; gate/up/down parameters all changed | + +The final run changed 16887 gate, 16272 up, and 14711 down values in each 131072-element parameter sample. + +## End-to-End Result + +Configuration: + +- Model: `Qwen3-30B-A3B-Instruct-2507` +- Full-FT BF16 expert path, TP=2 NUMA partitions +- Batch size 1, sequence length 1024, gradient accumulation 1 +- One RTX 5090, GPU 6 +- 15 steps; steps 8-15 used for the stable comparison +- C++ staged profiler enabled; Torch trace disabled + +| Stage | Before | After | Speedup | +|---|---:|---:|---:| +| End-to-end step | 45.204 s | 5.102 s | 8.86x | +| Throughput | 22.653 token/s | 200.701 token/s | 8.86x | +| Forward | 0.704 s | 0.668 s | 1.05x | +| Backward | 38.595 s | 3.250 s | 11.88x | +| Optimizer | 0.789 s | 0.687 s | 1.15x | +| Base-weight reload | 4.892 s | 0.272 s | 17.98x | +| TP0 base dW | 35.753 s | 0.933 s | 38.32x | + +The backward/forward ratio falls from 54.9 to 4.87. Stable direct packing itself takes 0.261 s/step; TP partition and cleanup stages are zero. Step 1 optimizer initialization and the first full reload in step 2 are cold-start outliers and are excluded from stable throughput. + +The output is retained at: + +```text +/mnt/sft_yyj_yyj/qwen3-30b-a3b-fullft/outputs/20260717_070714-bench-1024-15step +``` + +It occupies 16 MiB because the large Torch trace was disabled. GPU 6 peaked at 31886 MiB. Its sampled utilization averaged 1.84% over model load and training, confirming that the remaining critical path is still mostly CPU-side. + +## Detailed Profile + +Stable steps 8-15, TP0 critical path: + +| Stage | Time per step | +|---|---:| +| TP backward total | 1.447 s | +| Base dW wall time | 0.933 s | +| TP wrapper buffer clear | 0.352 s | +| Initial CPU MoE forward | 0.344 s | +| Checkpoint recompute CPU MoE forward | 0.505 s | +| Backward weight repack | 0.493 s | + +Worker CPU times are summed across the NUMA pool and therefore are not additive wall time: + +| dW worker stage | Summed CPU time per step | +|---|---:| +| Pack A | 11.248 s | +| Pack B | 9.933 s | +| Gate/up kernel | 3.614 s | +| Down kernel | 1.881 s | +| BF16 store | 1.960 s | + +Packing is about 74% of the measured dW worker CPU time. It is now a better optimization target than replacing the BF16 microkernel. + +## AMX Regression + +The same tests were built on an Intel Xeon Platinum 8488C with AMX enabled. The binary contains `tileloadd` and `tdpbf16ps`, and all dW and repack cases pass. + +The optional benchmark in `test_bf16_dweight --benchmark` compares the previous direct AMX tile loop with the common driver at route K=1024: + +| Path | Median time | +|---|---:| +| Legacy tile loop | 1285.9 ns | +| Common driver | 1295.1 ns | +| Ratio | 1.0072 | + +The measured AMX overhead is 0.72%, below the 5% regression limit. + +## Next Targets + +1. Reduce route-panel packing cost through wider transpose/copy primitives and panel reuse across adjacent output tiles. +2. Avoid clearing inactive or already-overwritten gradient regions in the TP wrapper. +3. Separate unavoidable backward repack from overlap gaps and remove only the exposed wall-clock part. +4. Revisit checkpoint recompute after the CPU dW and buffer-clear costs are lower. + diff --git a/doc/en/SFT/README.md b/doc/en/SFT/README.md index 324fb925d..d1a6068ec 100644 --- a/doc/en/SFT/README.md +++ b/doc/en/SFT/README.md @@ -4,5 +4,6 @@ - [Fine-Tuning User Guide](./KTransformers-Fine-Tuning_User-Guide.md) - [KT-FT Fine-Tuning and Inference Loop](./Qwen3.5-SGLang-LoRA-Serving.md) - [Developer Technical Notes](./KTransformers-Fine-Tuning_Developer-Technical-Notes.md) +- [Qwen3-30B-A3B BF16 Full-FT Performance](./Qwen3-30B-A3B-Full-FT-BF16-Performance.md) - [DPO Tutorial](./DPO_tutorial.md) - [Injection Tutorial](./injection_tutorial.md) From 1e95053b15b32e6db8193fd852d62d051c6e7ef5 Mon Sep 17 00:00:00 2001 From: yyj Date: Fri, 17 Jul 2026 15:23:48 +0800 Subject: [PATCH 17/20] [docs](sft): remove Qwen3 Full-FT performance report --- doc/SUMMARY.md | 1 - .../Qwen3-30B-A3B-Full-FT-BF16-Performance.md | 113 ------------------ doc/en/SFT/README.md | 1 - 3 files changed, 115 deletions(-) delete mode 100644 doc/en/SFT/Qwen3-30B-A3B-Full-FT-BF16-Performance.md diff --git a/doc/SUMMARY.md b/doc/SUMMARY.md index 7944d71d5..d19f29ca3 100644 --- a/doc/SUMMARY.md +++ b/doc/SUMMARY.md @@ -11,7 +11,6 @@ - [KT-FT 微调推理闭环](zh/Qwen3.5-SGLang-LoRA-Serving_zh.md) - [Injection Tutorial](en/SFT/injection_tutorial.md) - [kt-sft developer tech notes](en/SFT/KTransformers-Fine-Tuning_Developer-Technical-Notes.md) - - [Qwen3-30B-A3B BF16 Full-FT Performance](en/SFT/Qwen3-30B-A3B-Full-FT-BF16-Performance.md) - [DPO tutorial](en/SFT/DPO_tutorial.md) diff --git a/doc/en/SFT/Qwen3-30B-A3B-Full-FT-BF16-Performance.md b/doc/en/SFT/Qwen3-30B-A3B-Full-FT-BF16-Performance.md deleted file mode 100644 index 8b1d68c54..000000000 --- a/doc/en/SFT/Qwen3-30B-A3B-Full-FT-BF16-Performance.md +++ /dev/null @@ -1,113 +0,0 @@ -# Qwen3-30B-A3B BF16 Full-FT Performance - -## Summary - -This work removes two CPU bottlenecks in raw-BF16 routed-expert full-parameter training: - -1. Expert weight gradients now aggregate all routed rows into BF16 mat-mat operations. The common driver dispatches to AVX512-BF16 on non-AMX CPUs and to the existing AMX microkernel on AMX CPUs. -2. Updated full BF16 expert weights are packed directly into persistent TP-local forward buffers. The reload no longer allocates, copies, and frees six temporary TP partitions per layer and step. - -On a dual-socket AMD EPYC 9355 with one RTX 5090, Qwen3-30B-A3B at batch 1 and sequence length 1024 improves from 22.653 to 200.701 token/s. The comparison uses steps 8-15 of two 15-step runs. - -## Implementation - -The main implementation is commit `34d2102` on `fullft-development`: - -- `BF16DWeightKernel` packs route-major activation and gradient panels into the existing raw BF16 `BufferA` and `BufferB` layouts. -- Work is split by `expert x intermediate tile x {gate/up, down}` and scheduled through the NUMA-local work-stealing pool. -- Gate and up share the packed input panel. Down reuses its packed intermediate panel. -- FP32 accumulation and BF16 gradient storage are unchanged. -- `BufferB::from_mat_strided()` packs a TP slice directly from the full `[E,F,H]` or `[E,H,F]` parameter tensor. -- The old scalar fallback remains available for non-BF16 kernels. - -The staged profiler distinguishes wall-clock stages from summed worker CPU time. Worker stages use the `backward.base_weight_grad.worker_cpu.*` namespace and intentionally have no wall-clock parent percentage. - -## Correctness - -| Check | Result | -|---|---| -| AVX512-BF16 dW routes `1,31,32,33,65,1792,1825` and tail shapes | Exact match after BF16 store | -| Contiguous, transposed, and strided raw-BF16 repack | Pass | -| Synthetic full-FT gradients, TP1 and TP2, qlen `8,31,32,33,65` | 30/30 pass | -| Real HF layer-0 expert-0 gradients, TP1 and TP2 | Relative L2 `0.00355-0.00378`, cosine approximately 1 | -| 15-step end-to-end training | Finite loss and grad norm; gate/up/down parameters all changed | - -The final run changed 16887 gate, 16272 up, and 14711 down values in each 131072-element parameter sample. - -## End-to-End Result - -Configuration: - -- Model: `Qwen3-30B-A3B-Instruct-2507` -- Full-FT BF16 expert path, TP=2 NUMA partitions -- Batch size 1, sequence length 1024, gradient accumulation 1 -- One RTX 5090, GPU 6 -- 15 steps; steps 8-15 used for the stable comparison -- C++ staged profiler enabled; Torch trace disabled - -| Stage | Before | After | Speedup | -|---|---:|---:|---:| -| End-to-end step | 45.204 s | 5.102 s | 8.86x | -| Throughput | 22.653 token/s | 200.701 token/s | 8.86x | -| Forward | 0.704 s | 0.668 s | 1.05x | -| Backward | 38.595 s | 3.250 s | 11.88x | -| Optimizer | 0.789 s | 0.687 s | 1.15x | -| Base-weight reload | 4.892 s | 0.272 s | 17.98x | -| TP0 base dW | 35.753 s | 0.933 s | 38.32x | - -The backward/forward ratio falls from 54.9 to 4.87. Stable direct packing itself takes 0.261 s/step; TP partition and cleanup stages are zero. Step 1 optimizer initialization and the first full reload in step 2 are cold-start outliers and are excluded from stable throughput. - -The output is retained at: - -```text -/mnt/sft_yyj_yyj/qwen3-30b-a3b-fullft/outputs/20260717_070714-bench-1024-15step -``` - -It occupies 16 MiB because the large Torch trace was disabled. GPU 6 peaked at 31886 MiB. Its sampled utilization averaged 1.84% over model load and training, confirming that the remaining critical path is still mostly CPU-side. - -## Detailed Profile - -Stable steps 8-15, TP0 critical path: - -| Stage | Time per step | -|---|---:| -| TP backward total | 1.447 s | -| Base dW wall time | 0.933 s | -| TP wrapper buffer clear | 0.352 s | -| Initial CPU MoE forward | 0.344 s | -| Checkpoint recompute CPU MoE forward | 0.505 s | -| Backward weight repack | 0.493 s | - -Worker CPU times are summed across the NUMA pool and therefore are not additive wall time: - -| dW worker stage | Summed CPU time per step | -|---|---:| -| Pack A | 11.248 s | -| Pack B | 9.933 s | -| Gate/up kernel | 3.614 s | -| Down kernel | 1.881 s | -| BF16 store | 1.960 s | - -Packing is about 74% of the measured dW worker CPU time. It is now a better optimization target than replacing the BF16 microkernel. - -## AMX Regression - -The same tests were built on an Intel Xeon Platinum 8488C with AMX enabled. The binary contains `tileloadd` and `tdpbf16ps`, and all dW and repack cases pass. - -The optional benchmark in `test_bf16_dweight --benchmark` compares the previous direct AMX tile loop with the common driver at route K=1024: - -| Path | Median time | -|---|---:| -| Legacy tile loop | 1285.9 ns | -| Common driver | 1295.1 ns | -| Ratio | 1.0072 | - -The measured AMX overhead is 0.72%, below the 5% regression limit. - -## Next Targets - -1. Reduce route-panel packing cost through wider transpose/copy primitives and panel reuse across adjacent output tiles. -2. Avoid clearing inactive or already-overwritten gradient regions in the TP wrapper. -3. Separate unavoidable backward repack from overlap gaps and remove only the exposed wall-clock part. -4. Revisit checkpoint recompute after the CPU dW and buffer-clear costs are lower. - diff --git a/doc/en/SFT/README.md b/doc/en/SFT/README.md index d1a6068ec..324fb925d 100644 --- a/doc/en/SFT/README.md +++ b/doc/en/SFT/README.md @@ -4,6 +4,5 @@ - [Fine-Tuning User Guide](./KTransformers-Fine-Tuning_User-Guide.md) - [KT-FT Fine-Tuning and Inference Loop](./Qwen3.5-SGLang-LoRA-Serving.md) - [Developer Technical Notes](./KTransformers-Fine-Tuning_Developer-Technical-Notes.md) -- [Qwen3-30B-A3B BF16 Full-FT Performance](./Qwen3-30B-A3B-Full-FT-BF16-Performance.md) - [DPO Tutorial](./DPO_tutorial.md) - [Injection Tutorial](./injection_tutorial.md) From 66f5f15a7947f5ad03a71eec514c408148d4cc77 Mon Sep 17 00:00:00 2001 From: yyj Date: Fri, 17 Jul 2026 18:34:46 +0800 Subject: [PATCH 18/20] [perf](kt-kernel): reduce BF16 Full-FT checkpoint overhead Retain the first CPU MoE forward state across non-reentrant checkpoint recomputation, write BF16 activations directly into the backward cache, and reduce dWeight packing and gradient-clear traffic. Extend staged profiling and cover checkpoint reuse plus AMX/AVX dWeight paths. --- kt-kernel/operators/amx/la/bf16_dweight.hpp | 75 +++++- kt-kernel/operators/amx/moe_base.hpp | 10 +- kt-kernel/operators/amx/sft_moe.hpp | 218 ++++++++++++++++-- .../operators/amx/test/test_bf16_dweight.cpp | 135 ++++++++++- kt-kernel/operators/moe-sft-tp.hpp | 60 ++++- kt-kernel/operators/sft_profile.hpp | 20 ++ kt-kernel/python/sft/autograd.py | 44 +++- kt-kernel/python/sft/base.py | 28 +++ kt-kernel/python/sft/layer.py | 95 ++++---- kt-kernel/python/sft/profiler.py | 8 +- kt-kernel/python/sft/wrapper.py | 13 +- .../per_commit/test_sft_checkpoint_reuse.py | 113 +++++++++ .../test/per_commit/test_sft_profiler.py | 2 + 13 files changed, 720 insertions(+), 101 deletions(-) create mode 100644 kt-kernel/test/per_commit/test_sft_checkpoint_reuse.py diff --git a/kt-kernel/operators/amx/la/bf16_dweight.hpp b/kt-kernel/operators/amx/la/bf16_dweight.hpp index 02756fd9a..190917c59 100644 --- a/kt-kernel/operators/amx/la/bf16_dweight.hpp +++ b/kt-kernel/operators/amx/la/bf16_dweight.hpp @@ -17,6 +17,10 @@ struct BF16DWeightTimings { uint64_t pack_a_calls = 0; uint64_t pack_b_ns = 0; uint64_t pack_b_calls = 0; + uint64_t panel_input_ns = 0; + uint64_t panel_input_calls = 0; + uint64_t panel_grad_output_ns = 0; + uint64_t panel_grad_output_calls = 0; uint64_t kernel_gate_up_ns = 0; uint64_t kernel_gate_up_calls = 0; uint64_t kernel_down_ns = 0; @@ -101,10 +105,12 @@ class BF16DWeightKernel { static void configure_worker() { Kernel::config(); } static void pack_a_transposed(BufferA& destination, const ggml_bf16_t* source, int source_stride, int source_column, - int row_count, int routes) { + int row_count, int routes, int destination_row = 0) { + assert(destination_row >= 0 && destination_row + row_count <= destination.max_m); + assert(destination_row % M_STEP == 0 && row_count <= M_STEP); const int k = destination.k; for (int k_begin = 0; k_begin < k; k_begin += K_STEP) { - ggml_bf16_t* tile = destination.get_submat(M_STEP, k, 0, k_begin); + ggml_bf16_t* tile = destination.get_submat(destination.max_m, k, destination_row, k_begin); std::memset(tile, 0, M_STEP * K_STEP * sizeof(ggml_bf16_t)); const int valid_k = std::min(K_STEP, routes - k_begin); if (valid_k <= 0) continue; @@ -117,10 +123,12 @@ class BF16DWeightKernel { } static void pack_b_transposed(BufferB& destination, const ggml_bf16_t* source, int source_stride, int source_column, - int row_count, int routes) { + int row_count, int routes, int destination_row = 0) { + assert(destination_row >= 0 && destination_row + row_count <= destination.n); + assert(destination_row % N_STEP == 0 && row_count <= N_STEP); const int k = destination.k; for (int k_begin = 0; k_begin < k; k_begin += K_STEP) { - ggml_bf16_t* tile = destination.get_submat(N_STEP, k, 0, k_begin); + ggml_bf16_t* tile = destination.get_submat(destination.n, k, destination_row, k_begin); std::memset(tile, 0, N_STEP * K_STEP * sizeof(ggml_bf16_t)); const int valid_k = std::min(K_STEP, routes - k_begin); if (valid_k > 0) { @@ -135,13 +143,60 @@ class BF16DWeightKernel { } } - static void multiply(int padded_k, float* destination, BufferA& a, BufferB& b) { - for (int k_block_begin = 0; k_block_begin < padded_k; k_block_begin += Kernel::K_BLOCK) { - if constexpr (AMX_AVAILABLE) { - Kernel::amx_kernel(M_STEP, N_STEP, padded_k, 0, 0, k_block_begin, destination, &a, &b); - } else { - Kernel::avx_kernel_4(M_STEP, N_STEP, padded_k, 0, 0, k_block_begin, destination, &a, &b); + private: + template + static inline __attribute__((always_inline)) void multiply_avx_rows( + int padded_k, float* destination, BufferA& a, BufferB& b, int m_begin, int n_begin, int row_offset) { + static_assert(ROWS > 0 && ROWS <= M_STEP); + __m512 accum_lo[ROWS]; + __m512 accum_hi[ROWS]; + +#pragma GCC unroll 12 + for (int row = 0; row < ROWS; ++row) { + accum_lo[row] = _mm512_setzero_ps(); + accum_hi[row] = _mm512_setzero_ps(); + } + + for (int k_begin = 0; k_begin < padded_k; k_begin += K_STEP) { + const auto* a_pairs = reinterpret_cast(a.get_submat(a.max_m, padded_k, m_begin, k_begin)); + const auto* b_vectors = reinterpret_cast(b.get_submat(b.n, padded_k, n_begin, k_begin)); + + for (int k_pair = 0; k_pair < K_STEP / 2; ++k_pair) { + const __m512bh b_lo = b_vectors[k_pair]; + const __m512bh b_hi = b_vectors[Kernel::TILE_N + k_pair]; +#pragma GCC unroll 12 + for (int row = 0; row < ROWS; ++row) { + const __m512bh a_pair = + reinterpret_cast<__m512bh>(_mm512_set1_epi32(a_pairs[(row_offset + row) * (K_STEP / 2) + k_pair])); + accum_lo[row] = _mm512_dpbf16_ps(accum_lo[row], a_pair, b_lo); + accum_hi[row] = _mm512_dpbf16_ps(accum_hi[row], a_pair, b_hi); + } + } + } + +#pragma GCC unroll 12 + for (int row = 0; row < ROWS; ++row) { + _mm512_store_ps(destination + (row_offset + row) * N_STEP, accum_lo[row]); + _mm512_store_ps(destination + (row_offset + row) * N_STEP + Kernel::TILE_N, accum_hi[row]); + } + } + + public: + + static void multiply(int padded_k, float* destination, BufferA& a, BufferB& b, int m_begin = 0, int n_begin = 0) { + assert(m_begin >= 0 && m_begin + M_STEP <= a.max_m); + assert(n_begin >= 0 && n_begin + N_STEP <= b.n); + assert(m_begin % M_STEP == 0 && n_begin % N_STEP == 0); + if constexpr (AMX_AVAILABLE) { + for (int k_block_begin = 0; k_block_begin < padded_k; k_block_begin += Kernel::K_BLOCK) { + Kernel::amx_kernel(a.max_m, b.n, padded_k, m_begin, n_begin, k_block_begin, destination, &a, &b); } + } else { + // Keep the complete 32x32 FP32 output tile in registers. The generic forward + // kernel spans all 32 rows at once and spills its 64 accumulators on AVX512. + multiply_avx_rows<12>(padded_k, destination, a, b, m_begin, n_begin, 0); + multiply_avx_rows<12>(padded_k, destination, a, b, m_begin, n_begin, 12); + multiply_avx_rows<8>(padded_k, destination, a, b, m_begin, n_begin, 24); } } diff --git a/kt-kernel/operators/amx/moe_base.hpp b/kt-kernel/operators/amx/moe_base.hpp index 48c4da9ee..7329ff01d 100644 --- a/kt-kernel/operators/amx/moe_base.hpp +++ b/kt-kernel/operators/amx/moe_base.hpp @@ -673,21 +673,27 @@ class AMX_MOE_BASE { } void apply_activation(int activated_expert, int nth, int qlen) { + apply_activation_to(activated_expert, nth, qlen, m_local_gate_output_ptr_); + } + + void apply_activation_to(int activated_expert, int nth, int qlen, + const std::vector& destination_ptrs) { auto pool = config_.pool->get_subpool(tp_part_idx); - auto fn = [this, nth](int task_id) { + auto fn = [this, nth, &destination_ptrs](int task_id) { int expert_idx = m_expert_id_map_[task_id / nth]; int ith = task_id % nth; auto [n_start, n_end] = T::split_range_n(config_.intermediate_size, ith, nth); for (int i = 0; i < m_local_num_[expert_idx]; i++) { ggml_bf16_t* gate_output_ptr = &m_local_gate_output_ptr_[expert_idx][i * config_.intermediate_size]; ggml_bf16_t* up_output_ptr = &m_local_up_output_ptr_[expert_idx][i * config_.intermediate_size]; + ggml_bf16_t* destination_ptr = &destination_ptrs[expert_idx][i * config_.intermediate_size]; for (int j = n_start; j < n_end; j += 32) { __m512 gate_val0, gate_val1, up_val0, up_val1; avx512_32xbf16_to_32xfp32((__m512i*)(gate_output_ptr + j), &gate_val0, &gate_val1); avx512_32xbf16_to_32xfp32((__m512i*)(up_output_ptr + j), &up_val0, &up_val1); __m512 result0 = amx::act_fn(gate_val0, up_val0, config_.swiglu_limit, config_.swiglu_alpha); __m512 result1 = amx::act_fn(gate_val1, up_val1, config_.swiglu_limit, config_.swiglu_alpha); - avx512_32xfp32_to_32xbf16(&result0, &result1, (__m512i*)(gate_output_ptr + j)); + avx512_32xfp32_to_32xbf16(&result0, &result1, (__m512i*)(destination_ptr + j)); } } }; diff --git a/kt-kernel/operators/amx/sft_moe.hpp b/kt-kernel/operators/amx/sft_moe.hpp index 836f0ddac..c0aef51b0 100644 --- a/kt-kernel/operators/amx/sft_moe.hpp +++ b/kt-kernel/operators/amx/sft_moe.hpp @@ -474,10 +474,6 @@ class AMX_SFT_MOE_TP : public BaseMOE { // Last backward expert token distribution (for load balancing analysis) std::vector last_backward_expert_tokens_; - // Experts that had non-zero contributions in last backward (for selective zeroing) - std::vector last_backward_active_experts_; - bool grad_outputs_initialized_ = false; - // Cache buffer pools void* cache_input_pool_ = nullptr; void* cache_gate_output_pool_ = nullptr; @@ -512,6 +508,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { // Precomputed offsets for cache operations (avoid repeated heap allocation) std::vector cache_offsets_; + std::vector direct_cache_intermediate_ptrs_; // ===================================================== // AMX-optimized LoRA GEMM buffers (performance optimization) @@ -595,6 +592,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { void* backward_bc_pool_ = nullptr; void* grad_output_bf16_pool_ = nullptr; void* base_grad_output_bf16_pool_ = nullptr; + void* dweight_shared_panel_pool_ = nullptr; void* backward_pool_ = nullptr; size_t backward_pool_bytes_ = 0; @@ -603,6 +601,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { size_t backward_bc_pool_bytes_ = 0; size_t grad_output_bf16_pool_bytes_ = 0; size_t base_grad_output_bf16_pool_bytes_ = 0; + size_t dweight_shared_panel_pool_bytes_ = 0; // LoRA gradient computation pools (FP32, used in bwd_down_lora_precompute and grad computation) float* lora_grad_out_pool_ = nullptr; // [max_len * num_experts_per_tok * hidden_size] @@ -845,6 +844,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { required += round_up(backward_bc_pool_bytes_, kAmxAlignment); required += round_up(grad_output_bf16_pool_bytes_, kAmxAlignment); required += round_up(base_grad_output_bf16_pool_bytes_, kAmxAlignment); + required += round_up(dweight_shared_panel_pool_bytes_, kAmxAlignment); required += round_up(lora_grad_out_pool_bytes_, kAmxAlignment); required += round_up(lora_inter_proj_pool_bytes_, kAmxAlignment); required += round_up(lora_grad_times_b_pool_bytes_, kAmxAlignment); @@ -878,6 +878,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { assign(&backward_bc_pool_, backward_bc_pool_bytes_); assign(&grad_output_bf16_pool_, grad_output_bf16_pool_bytes_); assign(&base_grad_output_bf16_pool_, base_grad_output_bf16_pool_bytes_); + assign(&dweight_shared_panel_pool_, dweight_shared_panel_pool_bytes_); assign((void**)&lora_grad_out_pool_, lora_grad_out_pool_bytes_); assign((void**)&lora_inter_proj_pool_, lora_inter_proj_pool_bytes_); @@ -932,7 +933,21 @@ class AMX_SFT_MOE_TP : public BaseMOE { lora_scaling_ = rank > 0 ? alpha / rank : 0.0f; } - void set_full_weight_grad(bool enabled) { sft_config_.full_weight_grad = enabled; } + void set_full_weight_grad(bool enabled) { + sft_config_.full_weight_grad = enabled; + if constexpr (std::is_same_v) { + if (enabled) { + const size_t max_routes = static_cast(config_.max_len) * config_.num_experts_per_tok; + const size_t max_padded_routes = + ((max_routes + static_cast(config_.expert_num) * (T::K_STEP - 1) + T::K_STEP - 1) / T::K_STEP) * + T::K_STEP; + dweight_shared_panel_pool_bytes_ = + 2 * max_padded_routes * config_.hidden_size * sizeof(ggml_bf16_t); + } else { + dweight_shared_panel_pool_bytes_ = 0; + } + } + } /** * @brief SFT Forward pass with optional caching for backward. @@ -952,7 +967,8 @@ class AMX_SFT_MOE_TP : public BaseMOE { bool save_for_backward) { SFTProfileScope total_scope(profiler_, SFTProfileStage::FwdTotal); SFTProfileScope checkpoint_scope( - profiler_, save_for_backward ? SFTProfileStage::FwdRecomputeTotal : SFTProfileStage::FwdInitialTotal); + profiler_, save_for_backward && config_.share_cache_pool ? SFTProfileStage::FwdRecomputeTotal + : SFTProfileStage::FwdInitialTotal); auto stage_start = profiler_.start(); SFT_POOL_LOG("fwd_enter", config_.layer_idx, tp_part_idx, qlen, cache_stack_top_, forward_pool_bytes_, @@ -1124,6 +1140,10 @@ class AMX_SFT_MOE_TP : public BaseMOE { pool->do_work_stealing_job(count, nullptr, fn, nullptr); } }; + constexpr bool kSupportsDirectFullFTCache = std::is_same_v; + const bool use_direct_fullft_cache = + kSupportsDirectFullFTCache && save_for_backward && sft_config_.full_weight_grad && lora_rank_ == 0; + ForwardCache* direct_cache = nullptr; stage_start = profiler_.start(); direct_or_pool(qlen, [&](int i) { @@ -1159,6 +1179,24 @@ class AMX_SFT_MOE_TP : public BaseMOE { }); profiler_.record(SFTProfileStage::FwdInputPack, stage_start); + if constexpr (kSupportsDirectFullFTCache) { + if (use_direct_fullft_cache) { + stage_start = profiler_.start(); + ForwardCache& cache = (cache_stack_top_ > 0) ? cache_stack_[cache_stack_top_ - 1] : push_cache(); + prepare_cache_metadata(cache, qlen, k, expert_ids, weights, activated_expert); + bind_direct_cache_outputs(cache, activated_expert); + cache.valid = true; + direct_cache = &cache; + profiler_.record(SFTProfileStage::FwdCacheMetadata, stage_start); + + stage_start = profiler_.start(); + copy_input_to_cache(cache, input, qlen); + profiler_.record_bytes(SFTProfileStage::FwdCacheInput, + static_cast(qlen) * config_.hidden_size * sizeof(ggml_bf16_t)); + profiler_.record(SFTProfileStage::FwdCacheInput, stage_start); + } + } + // Step 5: Gate + Up GEMM (base projection) stage_start = profiler_.start(); int nth = T::recommended_nth(config_.intermediate_size); @@ -1229,8 +1267,10 @@ class AMX_SFT_MOE_TP : public BaseMOE { // overwrite it instead of pushing a new one. This keeps the cache // consistent with the current forward's buffer state (max_m, routing) // and avoids cache stack overflow from duplicate pushes. - ForwardCache& cache = (cache_stack_top_ > 0) ? cache_stack_[cache_stack_top_ - 1] : push_cache(); - save_to_cache(cache, qlen, k, expert_ids, weights, activated_expert, input); + ForwardCache& cache = direct_cache != nullptr + ? *direct_cache + : ((cache_stack_top_ > 0) ? cache_stack_[cache_stack_top_ - 1] : push_cache()); + if (direct_cache == nullptr) save_to_cache(cache, qlen, k, expert_ids, weights, activated_expert, input); // NaN Check: Forward Cache - input, gate_output, up_output if (is_nan_check_enabled()) { @@ -1272,7 +1312,15 @@ class AMX_SFT_MOE_TP : public BaseMOE { // Step 6: Activation (silu(gate) * up) stage_start = profiler_.start(); - { Base::apply_activation(activated_expert, nth, qlen); } + if (direct_cache != nullptr) { + Base::apply_activation_to(activated_expert, nth, qlen, direct_cache_intermediate_ptrs_); + for (int i = 0; i < activated_expert; ++i) { + const int expert_idx = m_expert_id_map_[i]; + m_local_gate_output_ptr_[expert_idx] = direct_cache_intermediate_ptrs_[expert_idx]; + } + } else { + Base::apply_activation(activated_expert, nth, qlen); + } profiler_.record(SFTProfileStage::FwdActivation, stage_start); // NaN Check: Step 6 - Activation output (silu(gate) * up) @@ -1293,7 +1341,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { stage_start = profiler_.start(); if (save_for_backward) { ForwardCache& cache = cache_stack_[cache_stack_top_ - 1]; // Get the cache we just pushed - save_intermediate_to_cache(cache, activated_expert); + if (direct_cache == nullptr) save_intermediate_to_cache(cache, activated_expert); // NaN Check: Forward Cache - intermediate_cache if (is_nan_check_enabled()) { @@ -1394,7 +1442,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { stage_start = profiler_.start(); if (save_for_backward) { ForwardCache& cache = cache_stack_[cache_stack_top_ - 1]; // Get the cache we just pushed - save_down_output_to_cache(cache, activated_expert); + if (direct_cache == nullptr) save_down_output_to_cache(cache, activated_expert); } profiler_.record(SFTProfileStage::FwdCacheDown, stage_start); @@ -2028,6 +2076,87 @@ class AMX_SFT_MOE_TP : public BaseMOE { auto pool = config_.pool->get_subpool(tp_part_idx); const bool profile_inner = profiler_.enabled(); + std::vector expert_padded_k(activated_expert); + std::vector panel_offsets(activated_expert); + size_t panel_elements = 0; + for (int expert_task = 0; expert_task < activated_expert; ++expert_task) { + const int expert_idx = cache.m_expert_id_map_cache[expert_task]; + panel_offsets[expert_task] = panel_elements; + expert_padded_k[expert_task] = DWeightKernel::padded_k(cache.m_local_num_cache[expert_idx]); + panel_elements += static_cast(H) * expert_padded_k[expert_task]; + } + const size_t required_panel_bytes = 2 * panel_elements * sizeof(ggml_bf16_t); + if (required_panel_bytes > dweight_shared_panel_pool_bytes_) { + throw std::runtime_error("BF16 dWeight shared panel pool is too small"); + } + auto* panel_base = static_cast(dweight_shared_panel_pool_); + auto* input_panels = panel_base; + auto* grad_output_panels = panel_base + panel_elements; + + // X and route-weighted dY are reused by every I tile. Pack each H tile once per expert, + // then keep the existing expert x I-tile compute schedule for load balancing. + stage_start = profiler_.start(); + const int panel_tasks_per_expert = h_tiles * 2; + const int panel_tasks = activated_expert * panel_tasks_per_expert; + if (panel_tasks > 0) { + pool->do_work_stealing_job( + panel_tasks, + [](int _) { amx::bf16_dweight_timings().reset(); }, + [&, h_tiles, panel_tasks_per_expert](int task_id) { + const int expert_task = task_id / panel_tasks_per_expert; + const int local_task = task_id % panel_tasks_per_expert; + const bool pack_grad_output = local_task >= h_tiles; + const int h_tile = local_task % h_tiles; + const int h_start = h_tile * TILE_N; + const int h_count = std::min(TILE_N, H - h_start); + const int expert_idx = cache.m_expert_id_map_cache[expert_task]; + const int routes = cache.m_local_num_cache[expert_idx]; + if (routes == 0) return; + + const int padded_k = expert_padded_k[expert_task]; + const size_t panel_offset = panel_offsets[expert_task]; + auto& timings = amx::bf16_dweight_timings(); + const auto profile_operation = [profile_inner](uint64_t& elapsed_ns, uint64_t& calls, + auto&& operation) { + if (!profile_inner) { + operation(); + return; + } + const auto begin = SFTProfiler::Clock::now(); + operation(); + elapsed_ns += static_cast( + std::chrono::duration_cast(SFTProfiler::Clock::now() - begin).count()); + calls++; + }; + + if (pack_grad_output) { + typename DWeightKernel::BufferA panel(H, padded_k, grad_output_panels + panel_offset); + profile_operation(timings.panel_grad_output_ns, timings.panel_grad_output_calls, [&] { + DWeightKernel::pack_a_transposed(panel, base_grad_output_bf16_ptr_[expert_idx], H, h_start, h_count, + routes, h_start); + }); + } else { + typename DWeightKernel::BufferB panel(H, padded_k, input_panels + panel_offset); + profile_operation(timings.panel_input_ns, timings.panel_input_calls, [&] { + DWeightKernel::pack_b_transposed(panel, m_local_input_ptr_[expert_idx], H, h_start, h_count, routes, + h_start); + }); + } + }, + [this](int _) { + auto& timings = amx::bf16_dweight_timings(); + profiler_.record_ns(SFTProfileStage::BwdBaseWeightGradPanelInput, timings.panel_input_ns, + timings.panel_input_calls); + profiler_.record_ns(SFTProfileStage::BwdBaseWeightGradPanelGradOutput, timings.panel_grad_output_ns, + timings.panel_grad_output_calls); + }); + } + profiler_.record(SFTProfileStage::BwdBaseWeightGradPanelPack, stage_start); + profiler_.record_bytes(SFTProfileStage::BwdBaseWeightGradPanelInput, + panel_elements * sizeof(ggml_bf16_t)); + profiler_.record_bytes(SFTProfileStage::BwdBaseWeightGradPanelGradOutput, + panel_elements * sizeof(ggml_bf16_t)); + stage_start = profiler_.start(); if (total_tasks > 0) { pool->do_work_stealing_job( @@ -2050,11 +2179,14 @@ class AMX_SFT_MOE_TP : public BaseMOE { const int padded_k = DWeightKernel::padded_k(routes); const int i_start = i_tile * TILE_M; const int i_count = std::min(TILE_M, I - i_start); + const size_t panel_offset = panel_offsets[expert_task]; auto& scratch = amx::bf16_dweight_scratch(); auto& timings = amx::bf16_dweight_timings(); typename DWeightKernel::BufferA a0(TILE_M, padded_k, scratch.a0()); typename DWeightKernel::BufferA a1(TILE_M, padded_k, scratch.a1()); typename DWeightKernel::BufferB b(TILE_N, padded_k, scratch.b()); + typename DWeightKernel::BufferA shared_grad_output(H, padded_k, grad_output_panels + panel_offset); + typename DWeightKernel::BufferB shared_input(H, padded_k, input_panels + panel_offset); auto profile_operation = [profile_inner](uint64_t& elapsed_ns, uint64_t& calls, auto&& operation) { if (!profile_inner) { @@ -2069,7 +2201,6 @@ class AMX_SFT_MOE_TP : public BaseMOE { }; if (do_down) { - const ggml_bf16_t* grad_output = base_grad_output_bf16_ptr_[expert_idx]; const ggml_bf16_t* intermediate = cache.intermediate_cache + pos_start * I; profile_operation(timings.pack_b_ns, timings.pack_b_calls, [&] { DWeightKernel::pack_b_transposed(b, intermediate, I, i_start, i_count, routes); @@ -2079,11 +2210,11 @@ class AMX_SFT_MOE_TP : public BaseMOE { for (int h_tile = 0; h_tile < h_tiles; ++h_tile) { const int h_start = h_tile * TILE_M; const int h_count = std::min(TILE_M, H - h_start); - profile_operation(timings.pack_a_ns, timings.pack_a_calls, [&] { - DWeightKernel::pack_a_transposed(a0, grad_output, H, h_start, h_count, routes); - }); profile_operation(timings.kernel_down_ns, timings.kernel_down_calls, - [&] { DWeightKernel::multiply(padded_k, scratch.c0(), a0, b); }); + [&] { + DWeightKernel::multiply(padded_k, scratch.c0(), shared_grad_output, b, h_start, + 0); + }); profile_operation(timings.store_ns, timings.store_calls, [&] { DWeightKernel::store_bf16(scratch.c0(), down_dst + static_cast(h_start) * F + i_start, F, h_count, i_count); @@ -2099,17 +2230,14 @@ class AMX_SFT_MOE_TP : public BaseMOE { DWeightKernel::pack_a_transposed(a1, up_grad, I, i_start, i_count, routes); }); - const ggml_bf16_t* input = m_local_input_ptr_[expert_idx]; ggml_bf16_t* gate_dst = ggp + static_cast(expert_idx) * F * H; ggml_bf16_t* up_dst = gup_ptr + static_cast(expert_idx) * F * H; for (int h_tile = 0; h_tile < h_tiles; ++h_tile) { const int h_start = h_tile * TILE_N; const int h_count = std::min(TILE_N, H - h_start); - profile_operation(timings.pack_b_ns, timings.pack_b_calls, - [&] { DWeightKernel::pack_b_transposed(b, input, H, h_start, h_count, routes); }); profile_operation(timings.kernel_gate_up_ns, timings.kernel_gate_up_calls, [&] { - DWeightKernel::multiply(padded_k, scratch.c0(), a0, b); - DWeightKernel::multiply(padded_k, scratch.c1(), a1, b); + DWeightKernel::multiply(padded_k, scratch.c0(), a0, shared_input, 0, h_start); + DWeightKernel::multiply(padded_k, scratch.c1(), a1, shared_input, 0, h_start); }); profile_operation(timings.store_ns, timings.store_calls, [&] { DWeightKernel::store_bf16(scratch.c0(), gate_dst + static_cast(i_start) * H + h_start, H, @@ -3159,6 +3287,14 @@ class AMX_SFT_MOE_TP : public BaseMOE { cache_down_output_bytes_ = (size_t)max_cache_depth_ * ml * k_tok * H * sizeof(ggml_bf16_t); grad_buffer_bytes_ = ml * k_tok * I * sizeof(ggml_bf16_t); + if constexpr (std::is_same_v) { + if (sft_config_.full_weight_grad) { + const size_t max_padded_routes = + ((ml * k_tok + static_cast(config_.expert_num) * (K_STEP - 1) + K_STEP - 1) / K_STEP) * K_STEP; + dweight_shared_panel_pool_bytes_ = + 2 * max_padded_routes * H * sizeof(ggml_bf16_t); // X^T BufferB + route-weighted dY^T BufferA + } + } // ===================================================== // Calculate LoRA AMX buffer sizes @@ -3366,6 +3502,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { cache_stack_.resize(max_cache_depth_); // Preallocate cache offsets to avoid heap allocation in hot path cache_offsets_.resize(config_.expert_num + 1); + direct_cache_intermediate_ptrs_.resize(config_.expert_num, nullptr); for (int i = 0; i < max_cache_depth_; i++) { // Note: cache pointers (input_cache, gate_output_cache, etc.) are set in alloc_forward_buffers() cache_stack_[i].input_cache = nullptr; @@ -4393,10 +4530,8 @@ class AMX_SFT_MOE_TP : public BaseMOE { return cache_stack_[--cache_stack_top_]; } - void save_to_cache(ForwardCache& cache, int qlen, int k, const int64_t* expert_ids, const float* weights, - int activated_expert, const void* input) { - auto pool = config_.pool->get_subpool(tp_part_idx); - + void prepare_cache_metadata(ForwardCache& cache, int qlen, int k, const int64_t* expert_ids, const float* weights, + int activated_expert) { cache.qlen_cache = qlen; cache.k_cache = k; cache.activated_expert_cache = activated_expert; @@ -4422,6 +4557,39 @@ class AMX_SFT_MOE_TP : public BaseMOE { int expert_idx = m_expert_id_map_[i]; cache_offsets_[i + 1] = cache_offsets_[i] + m_local_num_[expert_idx]; } + } + + void copy_input_to_cache(ForwardCache& cache, const void* input, int qlen) { + auto pool = config_.pool->get_subpool(tp_part_idx); + const size_t total_bytes = static_cast(qlen) * config_.hidden_size * sizeof(ggml_bf16_t); + constexpr size_t kChunkBytes = 2 * 1024 * 1024; + const int chunks = static_cast((total_bytes + kChunkBytes - 1) / kChunkBytes); + pool->do_work_stealing_job( + chunks, nullptr, + [&, input, total_bytes](int chunk) { + const size_t offset = static_cast(chunk) * kChunkBytes; + const size_t bytes = std::min(kChunkBytes, total_bytes - offset); + std::memcpy(reinterpret_cast(cache.input_cache) + offset, + reinterpret_cast(input) + offset, bytes); + }, + nullptr); + } + + void bind_direct_cache_outputs(ForwardCache& cache, int activated_expert) { + for (int i = 0; i < activated_expert; ++i) { + const int expert_idx = m_expert_id_map_[i]; + const size_t offset = cache_offsets_[i]; + m_local_gate_output_ptr_[expert_idx] = cache.gate_output_cache + offset * config_.intermediate_size; + m_local_up_output_ptr_[expert_idx] = cache.up_output_cache + offset * config_.intermediate_size; + m_local_down_output_ptr_[expert_idx] = cache.down_output_cache + offset * config_.hidden_size; + direct_cache_intermediate_ptrs_[expert_idx] = cache.intermediate_cache + offset * config_.intermediate_size; + } + } + + void save_to_cache(ForwardCache& cache, int qlen, int k, const int64_t* expert_ids, const float* weights, + int activated_expert, const void* input) { + auto pool = config_.pool->get_subpool(tp_part_idx); + prepare_cache_metadata(cache, qlen, k, expert_ids, weights, activated_expert); // Parallel copy: input(1 task) + gate(N tasks) + up(N tasks) = 1 + 2N tasks // This parallelizes the ~1.8MB input copy that was previously serial diff --git a/kt-kernel/operators/amx/test/test_bf16_dweight.cpp b/kt-kernel/operators/amx/test/test_bf16_dweight.cpp index 2e845aade..5afca4f87 100644 --- a/kt-kernel/operators/amx/test/test_bf16_dweight.cpp +++ b/kt-kernel/operators/amx/test/test_bf16_dweight.cpp @@ -94,6 +94,70 @@ bool run_case(int routes, int rows, int columns) { return passed; } +bool run_shared_panel_case(int routes, int rows, int columns) { + const int padded_k = DWeightKernel::padded_k(routes); + const int padded_rows = (rows + Kernel::M_STEP - 1) / Kernel::M_STEP * Kernel::M_STEP; + const int padded_columns = (columns + Kernel::N_STEP - 1) / Kernel::N_STEP * Kernel::N_STEP; + std::vector lhs(static_cast(routes) * rows); + std::vector rhs(static_cast(routes) * columns); + std::vector actual(static_cast(rows) * columns); + fill_random(lhs, static_cast(routes * 41 + rows)); + fill_random(rhs, static_cast(routes * 43 + columns)); + + void* a_memory = alloc_buffer(Kernel::BufferA::required_size(padded_rows, padded_k)); + void* b_memory = alloc_buffer(Kernel::BufferB::required_size(padded_columns, padded_k)); + Kernel::BufferA a(padded_rows, padded_k, a_memory); + Kernel::BufferB b(padded_columns, padded_k, b_memory); + for (int row = 0; row < rows; row += Kernel::M_STEP) { + DWeightKernel::pack_a_transposed(a, lhs.data(), rows, row, std::min(Kernel::M_STEP, rows - row), routes, row); + } + for (int column = 0; column < columns; column += Kernel::N_STEP) { + DWeightKernel::pack_b_transposed(b, rhs.data(), columns, column, + std::min(Kernel::N_STEP, columns - column), routes, column); + } + + alignas(64) float accumulator[Kernel::M_STEP * Kernel::N_STEP]; + for (int row = 0; row < rows; row += Kernel::M_STEP) { + const int row_count = std::min(Kernel::M_STEP, rows - row); + for (int column = 0; column < columns; column += Kernel::N_STEP) { + const int column_count = std::min(Kernel::N_STEP, columns - column); + DWeightKernel::multiply(padded_k, accumulator, a, b, row, column); + DWeightKernel::store_bf16(accumulator, actual.data() + static_cast(row) * columns + column, columns, + row_count, column_count); + } + } + + double difference_sq = 0.0; + double expected_sq = 0.0; + double actual_sq = 0.0; + double dot = 0.0; + for (int row = 0; row < rows; ++row) { + for (int column = 0; column < columns; ++column) { + float reference = 0.0f; + for (int route = 0; route < routes; ++route) { + reference += GGML_BF16_TO_FP32(lhs[static_cast(route) * rows + row]) * + GGML_BF16_TO_FP32(rhs[static_cast(route) * columns + column]); + } + const float expected = GGML_BF16_TO_FP32(GGML_FP32_TO_BF16(reference)); + const float value = GGML_BF16_TO_FP32(actual[static_cast(row) * columns + column]); + const double difference = static_cast(value) - expected; + difference_sq += difference * difference; + expected_sq += static_cast(expected) * expected; + actual_sq += static_cast(value) * value; + dot += static_cast(expected) * value; + } + } + const double relative_l2 = std::sqrt(difference_sq / std::max(expected_sq, 1e-30)); + const double cosine = dot / std::sqrt(std::max(expected_sq * actual_sq, 1e-30)); + const bool passed = relative_l2 <= 0.01 && cosine >= 0.999; + std::printf("BF16 dWeight shared panel routes=%d shape=%dx%d: rel_l2=%.6e cosine=%.9f %s\n", routes, rows, + columns, relative_l2, cosine, passed ? "PASS" : "FAIL"); + + std::free(a_memory); + std::free(b_memory); + return passed; +} + bool run_amx_benchmark() { if constexpr (!amx::AMX_AVAILABLE) { std::printf("BF16 dWeight AMX benchmark: SKIP (AMX unavailable)\n"); @@ -165,16 +229,85 @@ bool run_amx_benchmark() { return passed; } +bool run_avx_benchmark() { + if constexpr (amx::AMX_AVAILABLE) { + std::printf("BF16 dWeight AVX benchmark: SKIP (AMX enabled)\n"); + return true; + } + + constexpr int routes = 64; + constexpr int iterations = 50000; + constexpr int rounds = 7; + const int padded_k = DWeightKernel::padded_k(routes); + std::vector lhs(static_cast(routes) * Kernel::M_STEP); + std::vector rhs(static_cast(routes) * Kernel::N_STEP); + fill_random(lhs, 20260717); + fill_random(rhs, 20260718); + + void* a_memory = alloc_buffer(Kernel::BufferA::required_size(Kernel::M_STEP, padded_k)); + void* b_memory = alloc_buffer(Kernel::BufferB::required_size(Kernel::N_STEP, padded_k)); + Kernel::BufferA a(Kernel::M_STEP, padded_k, a_memory); + Kernel::BufferB b(Kernel::N_STEP, padded_k, b_memory); + alignas(64) float accumulator[Kernel::M_STEP * Kernel::N_STEP]; + DWeightKernel::pack_a_transposed(a, lhs.data(), Kernel::M_STEP, 0, Kernel::M_STEP, routes); + DWeightKernel::pack_b_transposed(b, rhs.data(), Kernel::N_STEP, 0, Kernel::N_STEP, routes); + + auto generic_driver = [&] { + for (int k_block_begin = 0; k_block_begin < padded_k; k_block_begin += Kernel::K_BLOCK) { + Kernel::avx_kernel_4(Kernel::M_STEP, Kernel::N_STEP, padded_k, 0, 0, k_block_begin, accumulator, &a, &b); + } + }; + auto register_blocked_driver = [&] { DWeightKernel::multiply(padded_k, accumulator, a, b); }; + + for (int warmup = 0; warmup < 200; ++warmup) { + generic_driver(); + register_blocked_driver(); + } + auto measure = [&](auto&& operation) { + const auto begin = std::chrono::steady_clock::now(); + for (int iteration = 0; iteration < iterations; ++iteration) operation(); + return std::chrono::duration(std::chrono::steady_clock::now() - begin).count() / iterations; + }; + + std::vector generic_ns; + std::vector register_blocked_ns; + for (int round = 0; round < rounds; ++round) { + if (round % 2 == 0) { + generic_ns.push_back(measure(generic_driver)); + register_blocked_ns.push_back(measure(register_blocked_driver)); + } else { + register_blocked_ns.push_back(measure(register_blocked_driver)); + generic_ns.push_back(measure(generic_driver)); + } + } + std::sort(generic_ns.begin(), generic_ns.end()); + std::sort(register_blocked_ns.begin(), register_blocked_ns.end()); + const double generic_median = generic_ns[rounds / 2]; + const double register_blocked_median = register_blocked_ns[rounds / 2]; + const double ratio = register_blocked_median / generic_median; + const bool passed = ratio <= 1.05; + std::printf("BF16 dWeight AVX kernel routes=%d: generic=%.1f ns register_blocked=%.1f ns ratio=%.4f %s\n", routes, + generic_median, register_blocked_median, ratio, passed ? "PASS" : "FAIL"); + + std::free(a_memory); + std::free(b_memory); + return passed; +} + } // namespace int main(int argc, char** argv) { DWeightKernel::configure_worker(); - if (argc == 2 && std::strcmp(argv[1], "--benchmark") == 0) return run_amx_benchmark() ? 0 : 1; + if (argc == 2 && std::strcmp(argv[1], "--benchmark") == 0) { + return (run_amx_benchmark() && run_avx_benchmark()) ? 0 : 1; + } bool passed = true; for (int routes : {1, 31, 32, 33, 65, 1792, 1825}) { passed = run_case(routes, 32, 32) && passed; } passed = run_case(33, 17, 29) && passed; passed = run_case(65, 31, 7) && passed; + passed = run_shared_panel_case(33, 45, 77) && passed; + passed = run_shared_panel_case(65, 64, 96) && passed; return passed ? 0 : 1; } diff --git a/kt-kernel/operators/moe-sft-tp.hpp b/kt-kernel/operators/moe-sft-tp.hpp index 12ca17cb8..1468fe5b9 100644 --- a/kt-kernel/operators/moe-sft-tp.hpp +++ b/kt-kernel/operators/moe-sft-tp.hpp @@ -217,6 +217,14 @@ class TP_MOE_SFT : public TP_MOE { std::vector part_grad_input_; std::vector part_grad_weights_; + // Full-FT dWeight outputs persist across calls. Active experts are overwritten in full; + // only experts that became inactive need to be cleared after the first invocation. + std::vector last_base_grad_active_experts_; + bool base_grad_outputs_initialized_ = false; + void* last_grad_gate_proj_ = nullptr; + void* last_grad_up_proj_ = nullptr; + void* last_grad_down_proj_ = nullptr; + public: TP_MOE_SFT(const MOESFTConfig& config) : Base(static_cast(config)), sft_config(config) { printf("Creating TP_MOE_SFT layer %d\n", config.layer_idx); @@ -724,18 +732,42 @@ class TP_MOE_SFT : public TP_MOE { } } + size_t base_grad_clear_bytes = 0; + const bool can_selectively_clear_base_grad = + need_base_weight_grad && T::kSupportsDirectBf16Reload && sft_config.lora_rank == 0 && + config.num_gpu_experts == 0; if (need_base_weight_grad) { const size_t base_grad_bytes = (size_t)expert_num * full_intermediate_size * hidden_size * sizeof(ggml_bf16_t); - auto append_clear_segments = [&](void* ptr) { - auto* base = static_cast(ptr); - for (size_t off = 0; off < base_grad_bytes; off += kChunkBytes) { - size_t len = std::min(kChunkBytes, base_grad_bytes - off); + const size_t expert_grad_bytes = (size_t)full_intermediate_size * hidden_size * sizeof(ggml_bf16_t); + auto append_clear_segments = [&](void* ptr, size_t offset, size_t bytes) { + auto* base = static_cast(ptr) + offset; + for (size_t off = 0; off < bytes; off += kChunkBytes) { + size_t len = std::min(kChunkBytes, bytes - off); clear_segs.push_back(ClearSeg{base + off, len}); } }; - append_clear_segments(grad_gate_proj); - append_clear_segments(grad_up_proj); - append_clear_segments(grad_down_proj); + + const bool pointers_unchanged = last_grad_gate_proj_ == grad_gate_proj && last_grad_up_proj_ == grad_up_proj && + last_grad_down_proj_ == grad_down_proj; + if (!can_selectively_clear_base_grad || !base_grad_outputs_initialized_ || !pointers_unchanged) { + append_clear_segments(grad_gate_proj, 0, base_grad_bytes); + append_clear_segments(grad_up_proj, 0, base_grad_bytes); + append_clear_segments(grad_down_proj, 0, base_grad_bytes); + base_grad_clear_bytes = 3 * base_grad_bytes; + } else { + std::vector active_mask(expert_num, 0); + for (int expert_idx : active_expert_map) { + if (expert_idx >= 0 && expert_idx < expert_num) active_mask[expert_idx] = 1; + } + for (int expert_idx : last_base_grad_active_experts_) { + if (expert_idx < 0 || expert_idx >= expert_num || active_mask[expert_idx]) continue; + const size_t expert_offset = static_cast(expert_idx) * expert_grad_bytes; + append_clear_segments(grad_gate_proj, expert_offset, expert_grad_bytes); + append_clear_segments(grad_up_proj, expert_offset, expert_grad_bytes); + append_clear_segments(grad_down_proj, expert_offset, expert_grad_bytes); + base_grad_clear_bytes += 3 * expert_grad_bytes; + } + } } pool->do_work_stealing_job((int)clear_segs.size(), nullptr, @@ -744,6 +776,9 @@ class TP_MOE_SFT : public TP_MOE { std::memset(seg.ptr, 0, seg.len); }, nullptr); + size_t partial_clear_bytes = 0; + for (size_t bytes : clear_bytes) partial_clear_bytes += bytes; + profiler_.record_bytes(SFTProfileStage::TpBwdBufferClear, partial_clear_bytes + base_grad_clear_bytes); profiler_.record(SFTProfileStage::TpBwdBufferClear, stage_start); // Compute TP-slice pointers for copy-type direct writes @@ -788,6 +823,9 @@ class TP_MOE_SFT : public TP_MOE { throw std::runtime_error("TP intermediate_size slices do not cover the full intermediate_size"); } + // A failed dWeight computation must force a full clear on the next invocation. + if (can_selectively_clear_base_grad) base_grad_outputs_initialized_ = false; + // Run backward on each NUMA node stage_start = profiler_.start(); pool->dispense_backend()->do_numa_job([&](int numa_id) { @@ -804,6 +842,13 @@ class TP_MOE_SFT : public TP_MOE { tp_grad_up_proj[numa_id], tp_grad_down_proj[numa_id]); }); profiler_.record(SFTProfileStage::TpBwdNumaCompute, stage_start); + if (can_selectively_clear_base_grad) { + last_base_grad_active_experts_ = active_expert_map; + last_grad_gate_proj_ = grad_gate_proj; + last_grad_up_proj_ = grad_up_proj; + last_grad_down_proj_ = grad_down_proj; + base_grad_outputs_initialized_ = true; + } // // Collect per-thread timing from all NUMA subpools // for (int i = 0; i < tp_count; i++) { @@ -1197,6 +1242,7 @@ class TP_MOE_SFT : public TP_MOE { */ void wait_backward_repack() { if (repack_thread_.joinable()) { + SFTProfileScope profile_scope(profiler_, SFTProfileStage::BackwardRepackWait); repack_thread_.join(); } } diff --git a/kt-kernel/operators/sft_profile.hpp b/kt-kernel/operators/sft_profile.hpp index 07548e1bd..22dd6b1cf 100644 --- a/kt-kernel/operators/sft_profile.hpp +++ b/kt-kernel/operators/sft_profile.hpp @@ -23,6 +23,8 @@ enum class SFTProfileStage : uint8_t { FwdBufferSetup, FwdInputScatter, FwdInputPack, + FwdCacheMetadata, + FwdCacheInput, FwdGateUpBase, FwdGateUpLora, FwdCacheGateUp, @@ -51,6 +53,9 @@ enum class SFTProfileStage : uint8_t { BwdRouterGrad, BwdBaseWeightGrad, BwdBaseWeightGradOffsets, + BwdBaseWeightGradPanelPack, + BwdBaseWeightGradPanelInput, + BwdBaseWeightGradPanelGradOutput, BwdBaseWeightGradMatMat, BwdBaseWeightGradPackA, BwdBaseWeightGradPackB, @@ -73,6 +78,7 @@ enum class SFTProfileStage : uint8_t { TpBwdLoraMerge, TpBwdRouterGradMerge, BackwardRepack, + BackwardRepackWait, BaseWeightReload, BaseWeightReloadPartition, BaseWeightReloadForwardPack, @@ -92,6 +98,8 @@ inline constexpr std::array(SFTProfileStage::Co "forward.buffer_setup", "forward.input_scatter", "forward.input_pack", + "forward.cache_metadata", + "forward.cache_input", "forward.gate_up_base", "forward.gate_up_lora", "forward.cache_gate_up", @@ -118,6 +126,9 @@ inline constexpr std::array(SFTProfileStage::Co "backward.router_grad", "backward.base_weight_grad", "backward.base_weight_grad.offsets", + "backward.base_weight_grad.panel_pack", + "backward.base_weight_grad.worker_cpu.panel_input", + "backward.base_weight_grad.worker_cpu.panel_grad_output", "backward.base_weight_grad.matmat", "backward.base_weight_grad.worker_cpu.pack_a", "backward.base_weight_grad.worker_cpu.pack_b", @@ -138,6 +149,7 @@ inline constexpr std::array(SFTProfileStage::Co "tp.backward.lora_merge", "tp.backward.router_grad_merge", "weights.backward_repack", + "weights.backward_repack_wait", "weights.base_reload", "weights.base_reload.partition", "weights.base_reload.forward_pack", @@ -176,6 +188,11 @@ class SFTProfiler { calls_[idx].fetch_add(calls, std::memory_order_relaxed); } + void record_bytes(SFTProfileStage stage, uint64_t bytes) { + if (!enabled_) return; + bytes_[static_cast(stage)].fetch_add(bytes, std::memory_order_relaxed); + } + void record_workload(uint64_t tokens, uint64_t routed_rows, uint64_t active_experts) { if (!enabled_) return; tokens_.fetch_add(tokens, std::memory_order_relaxed); @@ -194,12 +211,14 @@ class SFTProfiler { const std::string stage_prefix = prefix + kSFTProfileStageNames[i] + "."; out[stage_prefix + "total_ns"] = static_cast(load_or_exchange(total_ns_[i], reset_after)); out[stage_prefix + "calls"] = static_cast(load_or_exchange(calls_[i], reset_after)); + out[stage_prefix + "bytes"] = static_cast(load_or_exchange(bytes_[i], reset_after)); } } void reset() { for (auto& value : total_ns_) value.store(0, std::memory_order_relaxed); for (auto& value : calls_) value.store(0, std::memory_order_relaxed); + for (auto& value : bytes_) value.store(0, std::memory_order_relaxed); workloads_.store(0, std::memory_order_relaxed); tokens_.store(0, std::memory_order_relaxed); routed_rows_.store(0, std::memory_order_relaxed); @@ -214,6 +233,7 @@ class SFTProfiler { bool enabled_; std::array, static_cast(SFTProfileStage::Count)> total_ns_{}; std::array, static_cast(SFTProfileStage::Count)> calls_{}; + std::array, static_cast(SFTProfileStage::Count)> bytes_{}; std::atomic workloads_{0}; std::atomic tokens_{0}; std::atomic routed_rows_{0}; diff --git a/kt-kernel/python/sft/autograd.py b/kt-kernel/python/sft/autograd.py index bb9ce7974..615429a7c 100644 --- a/kt-kernel/python/sft/autograd.py +++ b/kt-kernel/python/sft/autograd.py @@ -38,6 +38,8 @@ def forward( training: bool, train_lora: bool, all_qlens: list[int] | tuple[int, ...] | None, + cache_checkpoint_forward: bool = False, + reuse_cached_forward: bool = False, gate_proj_param: torch.Tensor | None = None, up_proj_param: torch.Tensor | None = None, down_proj_param: torch.Tensor | None = None, @@ -78,8 +80,17 @@ def forward( # Rank 0: sync CPU result and split by real lengths if rank == 0: - with torch.profiler.record_function("kt.sft.cpu_forward_sync"): - cpu_output = wrapper.sync_forward(output_device=original_device) + if reuse_cached_forward: + with torch.profiler.record_function("kt.sft.checkpoint_cached_cpu_moe"): + cpu_output = wrapper.get_checkpoint_output(total_qlen, output_device=original_device) + elif cache_checkpoint_forward: + with torch.profiler.record_function("kt.sft.cpu_forward_sync"): + cached_output = wrapper.sync_forward(output_device=None) + wrapper.cache_checkpoint_output(cached_output, total_qlen) + cpu_output = cached_output.to(device=original_device, non_blocking=True) + else: + with torch.profiler.record_function("kt.sft.cpu_forward_sync"): + cpu_output = wrapper.sync_forward(output_device=original_device) cpu_output = cpu_output.to(dtype=original_dtype).view(total_qlen, hidden_size) offsets = _qlen_offsets(all_qlens_list) scatter_list = [cpu_output[offsets[i] : offsets[i + 1]].contiguous() for i in range(world_size)] @@ -99,8 +110,17 @@ def forward( del output_flat elif wrapper is not None: # Single-GPU: sync directly - with torch.profiler.record_function("kt.sft.cpu_forward_sync"): - cpu_output = wrapper.sync_forward(output_device=original_device) + if reuse_cached_forward: + with torch.profiler.record_function("kt.sft.checkpoint_cached_cpu_moe"): + cpu_output = wrapper.get_checkpoint_output(qlen, output_device=original_device) + elif cache_checkpoint_forward: + with torch.profiler.record_function("kt.sft.cpu_forward_sync"): + cached_output = wrapper.sync_forward(output_device=None) + wrapper.cache_checkpoint_output(cached_output, qlen) + cpu_output = cached_output.to(device=original_device, non_blocking=True) + else: + with torch.profiler.record_function("kt.sft.cpu_forward_sync"): + cpu_output = wrapper.sync_forward(output_device=original_device) output = cpu_output.view(batch_size, seq_len, hidden_size).to(dtype=original_dtype) else: # Broadcast-only rank (no wrapper) @@ -241,12 +261,14 @@ def backward(ctx, grad_output: torch.Tensor): elif not ctx.use_broadcast: # ---- Single-GPU path ---- grad_output_flat = grad_output.view(qlen, hidden_size) - with torch.profiler.record_function("kt.sft.cpu_backward"): - backward_out = ctx.wrapper.backward( - grad_output_flat, - output_device=ctx.original_device, - ) - ctx.wrapper._kt_has_cached_forward = False + try: + with torch.profiler.record_function("kt.sft.cpu_backward"): + backward_out = ctx.wrapper.backward( + grad_output_flat, + output_device=ctx.original_device, + ) + finally: + ctx.wrapper.clear_checkpoint_output() if isinstance(backward_out, tuple) and len(backward_out) == 2: grad_input, grad_weights = backward_out elif isinstance(backward_out, tuple) and len(backward_out) == 3: @@ -290,6 +312,8 @@ def backward(ctx, grad_output: torch.Tensor): None, None, None, + None, + None, grad_gate_proj, grad_up_proj, grad_down_proj, diff --git a/kt-kernel/python/sft/base.py b/kt-kernel/python/sft/base.py index c04f4cc6f..d402dc3fe 100644 --- a/kt-kernel/python/sft/base.py +++ b/kt-kernel/python/sft/base.py @@ -194,6 +194,10 @@ def __init__( self._cache_depth: int = 0 self._is_skip_lora: bool = False self._base_weights_dirty: bool = False + self.reuse_checkpoint_forward: bool = False + self._kt_has_cached_forward: bool = False + self._checkpoint_output_cpu: Optional[torch.Tensor] = None + self._checkpoint_output_qlen: int = 0 self.moe = None @@ -347,6 +351,30 @@ def _return_output(self, buffer: KExpertsSFTBuffer, qlen: int, output_device: Op else: return buffer.output_cpu[:qlen].clone() + def cache_checkpoint_output(self, output_cpu: torch.Tensor, qlen: int) -> None: + if output_cpu.device.type != "cpu": + raise ValueError("checkpoint CPU expert output must reside on CPU") + if output_cpu.shape[0] < qlen: + raise ValueError(f"checkpoint output is shorter than qlen: {output_cpu.shape[0]} < {qlen}") + self._checkpoint_output_cpu = output_cpu[:qlen].contiguous() + self._checkpoint_output_qlen = qlen + self._kt_has_cached_forward = True + + def get_checkpoint_output(self, qlen: int, output_device: Optional[torch.device] = None) -> torch.Tensor: + if not self._kt_has_cached_forward or self._checkpoint_output_cpu is None: + raise RuntimeError("No cached checkpoint forward output is available.") + if qlen != self._checkpoint_output_qlen: + raise RuntimeError(f"Cached checkpoint qlen mismatch: cached={self._checkpoint_output_qlen}, requested={qlen}") + output = self._checkpoint_output_cpu + if output_device is not None: + return output.to(device=output_device, non_blocking=True) + return output + + def clear_checkpoint_output(self) -> None: + self._checkpoint_output_cpu = None + self._checkpoint_output_qlen = 0 + self._kt_has_cached_forward = False + def _return_grads(self, buffer: KExpertsSFTBuffer, qlen: int, output_device: Optional[torch.device]): if output_device is not None: grad_input = buffer.grad_input_cpu[:qlen].to(device=output_device, non_blocking=True) diff --git a/kt-kernel/python/sft/layer.py b/kt-kernel/python/sft/layer.py index 7db6398b0..4f0acd902 100644 --- a/kt-kernel/python/sft/layer.py +++ b/kt-kernel/python/sft/layer.py @@ -148,8 +148,20 @@ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: and (hidden_states.requires_grad or topk_weights.requires_grad or train_lora or full_weight_grad) ) use_autograd_path = save_for_backward + checkpoint_mode = _checkpoint_hook_mode() + reuse_checkpoint_forward = ( + not dist_on + and self.wrapper is not None + and getattr(self.wrapper, "reuse_checkpoint_forward", False) + ) + reuse_cached_forward = ( + reuse_checkpoint_forward + and checkpoint_mode == "recompute" + and getattr(self.wrapper, "_kt_has_cached_forward", False) + ) + cache_checkpoint_forward = reuse_checkpoint_forward and checkpoint_mode == "first_forward" save_for_backward_submit = use_autograd_path - if _checkpoint_hook_mode() == "first_forward": + if checkpoint_mode == "first_forward" and not cache_checkpoint_forward: save_for_backward_submit = False if train_lora and self._lora_pointers_dirty: @@ -168,6 +180,7 @@ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: topk_ids, topk_weights, save_for_backward_submit, + reuse_cached_forward, ) # Use KTMoEFunction whenever backward is needed so KT backward and LoRA @@ -200,6 +213,8 @@ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: save_for_backward, train_lora, all_qlens, + cache_checkpoint_forward, + reuse_cached_forward, # Base weight params for full mode gradient flow self.wrapper.gate_proj_buf if full_weight_grad and self.wrapper is not None else None, self.wrapper.up_proj_buf if full_weight_grad and self.wrapper is not None else None, @@ -377,6 +392,7 @@ def _submit_and_compute_gpu( topk_ids: torch.Tensor, topk_weights: torch.Tensor, save_for_backward: bool, + reuse_cached_forward: bool = False, ) -> tuple[torch.Tensor | None, list[int] | None]: import torch.distributed as dist @@ -396,43 +412,37 @@ def _submit_and_compute_gpu( raise RuntimeError(f"Rank {rank} qlen mismatch: local={qlen}, all_qlens[{rank}]={all_qlens[rank]}") total_qlen = sum(all_qlens) - hs_flat = hidden_states.view(qlen, self.hidden_size).contiguous() - expert_ids = topk_ids.view(qlen, self.moe_config.num_experts_per_tok).contiguous() - weights = topk_weights.view(qlen, self.moe_config.num_experts_per_tok).contiguous() - - submit_hs = hs_flat.detach() - submit_ids = expert_ids.detach() - submit_wts = weights.detach() + if not reuse_cached_forward: + hs_flat = hidden_states.view(qlen, self.hidden_size).contiguous() + expert_ids = topk_ids.view(qlen, self.moe_config.num_experts_per_tok).contiguous() + weights = topk_weights.view(qlen, self.moe_config.num_experts_per_tok).contiguous() - gathered_hs = _dist_gather_varlen_to_rank0( - submit_hs, - all_qlens=all_qlens, - rank=rank, - world_size=world_size, - ) - gathered_ids = _dist_gather_varlen_to_rank0( - submit_ids, - all_qlens=all_qlens, - rank=rank, - world_size=world_size, - ) - gathered_wts = _dist_gather_varlen_to_rank0( - submit_wts, - all_qlens=all_qlens, - rank=rank, - world_size=world_size, - ) - - if rank == 0: - all_hs = torch.cat(gathered_hs, dim=0) - all_ids = torch.cat(gathered_ids, dim=0) - all_wts = torch.cat(gathered_wts, dim=0) - self.wrapper.submit_forward( - all_hs, - all_ids, - all_wts, - save_for_backward=save_for_backward, + gathered_hs = _dist_gather_varlen_to_rank0( + hs_flat.detach(), + all_qlens=all_qlens, + rank=rank, + world_size=world_size, + ) + gathered_ids = _dist_gather_varlen_to_rank0( + expert_ids.detach(), + all_qlens=all_qlens, + rank=rank, + world_size=world_size, ) + gathered_wts = _dist_gather_varlen_to_rank0( + weights.detach(), + all_qlens=all_qlens, + rank=rank, + world_size=world_size, + ) + + if rank == 0: + self.wrapper.submit_forward( + torch.cat(gathered_hs, dim=0), + torch.cat(gathered_ids, dim=0), + torch.cat(gathered_wts, dim=0), + save_for_backward=save_for_backward, + ) # Keep shared/lora experts local to avoid qlen_max-style amplification. gpu_output = None @@ -456,12 +466,13 @@ def _submit_and_compute_gpu( submit_hs = input_flat.detach() submit_ids = expert_ids.detach() submit_wts = weights.detach() - self.wrapper.submit_forward( - submit_hs, - submit_ids, - submit_wts, - save_for_backward=save_for_backward, - ) + if not reuse_cached_forward: + self.wrapper.submit_forward( + submit_hs, + submit_ids, + submit_wts, + save_for_backward=save_for_backward, + ) # GPU compute: shared_experts + lora_experts gpu_output = None diff --git a/kt-kernel/python/sft/profiler.py b/kt-kernel/python/sft/profiler.py index e909166cd..b3436fdbb 100644 --- a/kt-kernel/python/sft/profiler.py +++ b/kt-kernel/python/sft/profiler.py @@ -90,7 +90,7 @@ def _parent_stage(stage: str) -> str | None: def _aggregate_rows(profile: dict[str, Any]) -> list[dict[str, float | str]]: totals: dict[tuple[str, str], dict[str, float]] = defaultdict( - lambda: {"total_ns": 0.0, "calls": 0.0, "tokens": 0.0} + lambda: {"total_ns": 0.0, "calls": 0.0, "tokens": 0.0, "bytes": 0.0} ) for raw in profile.get("layers", {}).values(): scope_tokens: dict[str, float] = {"wrapper": raw.get("wrapper.tokens", 0.0)} @@ -108,6 +108,7 @@ def _aggregate_rows(profile: dict[str, Any]) -> list[dict[str, float | str]]: row["total_ns"] += total_ns row["calls"] += calls row["tokens"] += scope_tokens.get(scope, 0.0) + row["bytes"] += raw.get(key[: -len("total_ns")] + "bytes", 0.0) rows: list[dict[str, float | str]] = [] for (scope, stage), values in totals.items(): @@ -123,6 +124,7 @@ def _aggregate_rows(profile: dict[str, Any]) -> list[dict[str, float | str]]: "total_ms": values["total_ns"] / 1e6, "avg_ms": values["total_ns"] / calls / 1e6 if calls else 0.0, "us_per_token": values["total_ns"] / tokens / 1e3 if tokens else 0.0, + "mib": values["bytes"] / (1024.0 * 1024.0), "parent_pct": values["total_ns"] / parent_ns * 100.0 if parent_ns else 0.0, } ) @@ -137,13 +139,13 @@ def format_kt_sft_profile(profile: dict[str, Any]) -> str: rows = _aggregate_rows(profile) header = ( f"{'scope':<9} {'stage':<34} {'calls':>7} {'total_ms':>11} " - f"{'avg_ms':>10} {'us/token':>10} {'parent%':>9}" + f"{'avg_ms':>10} {'us/token':>10} {'MiB':>10} {'parent%':>9}" ) lines = [header, "-" * len(header)] for row in rows: lines.append( f"{row['scope']:<9} {row['stage']:<34} {row['calls']:>7.0f} " f"{row['total_ms']:>11.3f} {row['avg_ms']:>10.3f} " - f"{row['us_per_token']:>10.3f} {row['parent_pct']:>8.1f}%" + f"{row['us_per_token']:>10.3f} {row['mib']:>10.2f} {row['parent_pct']:>8.1f}%" ) return "\n".join(lines) diff --git a/kt-kernel/python/sft/wrapper.py b/kt-kernel/python/sft/wrapper.py index bb5284353..4e3dd4614 100644 --- a/kt-kernel/python/sft/wrapper.py +++ b/kt-kernel/python/sft/wrapper.py @@ -353,7 +353,18 @@ def wrap_moe_layers_with_kt_wrapper(model: nn.Module, kt_plugin: Any) -> list[KT # Set share_backward_bb and share_cache_pool BEFORE load_weights (config is built during load) wrapper.share_backward_bb = cfg.kt_share_backward_bb - wrapper.share_cache_pool = cfg.kt_share_cache_pool + single_process = not dist.is_initialized() or dist.get_world_size() == 1 + reuse_checkpoint_forward = ( + single_process + and full_weight_grad + and lora_rank == 0 + and kt_method == "AMXBF16_SFT" + and os.environ.get("KT_REUSE_CHECKPOINT_FORWARD", "1") != "0" + ) + wrapper.reuse_checkpoint_forward = reuse_checkpoint_forward + # Reusing the first checkpoint forward requires each layer's C++ + # activations to remain valid until its backward invocation. + wrapper.share_cache_pool = False if reuse_checkpoint_forward else cfg.kt_share_cache_pool physical_to_logical_map = torch.arange(moe_config.expert_num, dtype=torch.int64, device="cpu") diff --git a/kt-kernel/test/per_commit/test_sft_checkpoint_reuse.py b/kt-kernel/test/per_commit/test_sft_checkpoint_reuse.py new file mode 100644 index 000000000..b3d2e3037 --- /dev/null +++ b/kt-kernel/test/per_commit/test_sft_checkpoint_reuse.py @@ -0,0 +1,113 @@ +# SPDX-License-Identifier: Apache-2.0 + +import torch +from torch.utils.checkpoint import checkpoint + +from kt_kernel.sft.autograd import KTMoEFunction +from kt_kernel.sft.dist_utils import _checkpoint_hook_mode + + +class _FakeWrapper: + def __init__(self): + self._full_weight_grad = False + self.share_backward_bb = False + self._kt_has_cached_forward = False + self._checkpoint_output_cpu = None + self.submit_calls = 0 + self.sync_calls = 0 + self.cached_output_calls = 0 + self.backward_calls = 0 + + def submit_forward(self, hidden_states, _expert_ids, weights, save_for_backward=True): + self.submit_calls += 1 + self.input = hidden_states.detach().clone() + self.weights = weights.detach().clone() + self.output = self.input * self.weights + assert save_for_backward + + def sync_forward(self, output_device=None): + self.sync_calls += 1 + output = self.output.clone() + return output if output_device is None else output.to(output_device) + + def cache_checkpoint_output(self, output, _qlen): + self._checkpoint_output_cpu = output + self._kt_has_cached_forward = True + + def get_checkpoint_output(self, _qlen, output_device=None): + self.cached_output_calls += 1 + output = self._checkpoint_output_cpu + return output if output_device is None else output.to(output_device) + + def clear_checkpoint_output(self): + self._checkpoint_output_cpu = None + self._kt_has_cached_forward = False + + def backward(self, grad_output, output_device=None): + self.backward_calls += 1 + grad_input = grad_output * self.weights + grad_weights = (grad_output * self.input).sum(dim=-1, keepdim=True) + if output_device is not None: + grad_input = grad_input.to(output_device) + grad_weights = grad_weights.to(output_device) + return grad_input, grad_weights + + +class _CheckpointedExpert(torch.nn.Module): + def __init__(self): + super().__init__() + self.wrapper = _FakeWrapper() + self.route_weights = torch.nn.Parameter(torch.tensor([[0.25], [0.5], [0.75]], dtype=torch.bfloat16)) + + def forward(self, hidden_states): + batch, seq_len, hidden_size = hidden_states.shape + qlen = batch * seq_len + mode = _checkpoint_hook_mode() + cache_forward = mode == "first_forward" + reuse_forward = mode == "recompute" and self.wrapper._kt_has_cached_forward + expert_ids = torch.zeros((batch, seq_len, 1), dtype=torch.int64) + + if not reuse_forward: + self.wrapper.submit_forward( + hidden_states.view(qlen, hidden_size), + expert_ids.view(qlen, 1), + self.route_weights, + save_for_backward=True, + ) + + return KTMoEFunction.apply( + hidden_states, + expert_ids, + self.route_weights, + self.wrapper, + hidden_states.new_empty(()), + hidden_size, + 1, + 0, + True, + False, + None, + cache_forward, + reuse_forward, + None, + None, + None, + ) + + +def test_non_reentrant_checkpoint_reuses_cpu_expert_forward_and_preserves_gradients(): + module = _CheckpointedExpert() + hidden_states = torch.arange(12, dtype=torch.float32).view(1, 3, 4).requires_grad_(True) + + output = checkpoint(module, hidden_states, use_reentrant=False) + output.sum().backward() + + expected_input_grad = module.route_weights.detach().float().view(1, 3, 1).expand_as(hidden_states) + expected_weight_grad = hidden_states.detach().sum(dim=-1).view(3, 1).to(torch.bfloat16) + torch.testing.assert_close(hidden_states.grad, expected_input_grad) + torch.testing.assert_close(module.route_weights.grad, expected_weight_grad) + assert module.wrapper.submit_calls == 1 + assert module.wrapper.sync_calls == 1 + assert module.wrapper.cached_output_calls == 1 + assert module.wrapper.backward_calls == 1 + assert not module.wrapper._kt_has_cached_forward diff --git a/kt-kernel/test/per_commit/test_sft_profiler.py b/kt-kernel/test/per_commit/test_sft_profiler.py index f411f59df..fc9a0d8b3 100644 --- a/kt-kernel/test/per_commit/test_sft_profiler.py +++ b/kt-kernel/test/per_commit/test_sft_profiler.py @@ -19,6 +19,7 @@ def get_profile_stats(self, reset=False): "wrapper.tp.forward.total.calls": 2, "wrapper.tp.forward.numa_compute.total_ns": 1_500_000, "wrapper.tp.forward.numa_compute.calls": 2, + "wrapper.tp.forward.numa_compute.bytes": 2 * 1024 * 1024, "tp.0.enabled": 1, "tp.0.tokens": 8, "tp.0.forward.total.total_ns": 1_400_000, @@ -57,6 +58,7 @@ def test_collect_and_format_profile(): assert "tp.forward.numa_compute" in output assert "forward.route" in output assert "75.0%" in output + assert "2.00" in output assert "10.0%" in output assert "40.0%" in output worker_store = next(line for line in output.splitlines() if "worker_cpu.store" in line) From 6eeffa77a45fa3cbf5c3e9e6e5b2a783627b3c3f Mon Sep 17 00:00:00 2001 From: yyj Date: Sat, 18 Jul 2026 22:49:50 +0800 Subject: [PATCH 19/20] [perf](kt-kernel): make SFT optimizer gradients authoritative Bind Full-FT and LoRA Parameter.grad directly to the KT-managed BF16 gradient buffers, avoiding PyTorch duplicate accumulation. Accumulate microbatch gradients in C++, lazily clear expert buffers between optimizer windows, preserve rank-0 distributed ownership, and add lifecycle and AMX dWeight coverage. --- kt-kernel/ext_bindings.cpp | 27 +- kt-kernel/operators/amx/la/bf16_dweight.hpp | 16 +- kt-kernel/operators/amx/sft_moe.hpp | 61 ++- .../operators/amx/test/test_bf16_dweight.cpp | 58 +++ kt-kernel/operators/common.hpp | 4 + kt-kernel/operators/moe-sft-tp.hpp | 380 ++++++++++++------ kt-kernel/operators/sft_profile.hpp | 6 + kt-kernel/python/sft/amx.py | 25 +- kt-kernel/python/sft/autograd.py | 176 +++++--- kt-kernel/python/sft/base.py | 298 +++++++++++++- kt-kernel/python/sft/dist_utils.py | 32 ++ kt-kernel/python/sft/layer.py | 47 ++- kt-kernel/python/sft/lora.py | 156 +++++-- kt-kernel/python/sft/wrapper.py | 62 ++- .../per_commit/test_sft_authoritative_grad.py | 361 +++++++++++++++++ 15 files changed, 1417 insertions(+), 292 deletions(-) create mode 100644 kt-kernel/test/per_commit/test_sft_authoritative_grad.py diff --git a/kt-kernel/ext_bindings.cpp b/kt-kernel/ext_bindings.cpp index 9ce751752..0c6942667 100644 --- a/kt-kernel/ext_bindings.cpp +++ b/kt-kernel/ext_bindings.cpp @@ -328,6 +328,8 @@ class MOESFTBindings { intptr_t grad_gate_proj; intptr_t grad_up_proj; intptr_t grad_down_proj; + bool accumulate_optimizer_grads; + float optimizer_grad_scale; }; static void inner(void* args) { Args* args_ = (Args*)args; @@ -335,7 +337,8 @@ class MOESFTBindings { args_->grad_gate_lora_a, args_->grad_gate_lora_b, args_->grad_up_lora_a, args_->grad_up_lora_b, args_->grad_down_lora_a, args_->grad_down_lora_b, args_->grad_weights, args_->grad_gate_proj, args_->grad_up_proj, - args_->grad_down_proj); + args_->grad_down_proj, args_->accumulate_optimizer_grads, + args_->optimizer_grad_scale); } static std::pair cpuinfer_interface(std::shared_ptr> moe, intptr_t grad_output, intptr_t grad_input, intptr_t grad_gate_lora_a, @@ -343,11 +346,14 @@ class MOESFTBindings { intptr_t grad_up_lora_b, intptr_t grad_down_lora_a, intptr_t grad_down_lora_b, intptr_t grad_weights, intptr_t grad_gate_proj, intptr_t grad_up_proj, - intptr_t grad_down_proj) { + intptr_t grad_down_proj, + bool accumulate_optimizer_grads = false, + float optimizer_grad_scale = 1.0f) { Args* args = new Args{nullptr, moe.get(), grad_output, grad_input, grad_gate_lora_a, grad_gate_lora_b, grad_up_lora_a, grad_up_lora_b, grad_down_lora_a, grad_down_lora_b, grad_weights, - grad_gate_proj, grad_up_proj, grad_down_proj}; + grad_gate_proj, grad_up_proj, grad_down_proj, + accumulate_optimizer_grads, optimizer_grad_scale}; return std::make_pair((intptr_t)&inner, (intptr_t)args); } }; @@ -398,12 +404,22 @@ void bind_moe_sft_module(py::module_& moe_module, const char* name) { .def("warm_up_task", &MoeBindings::WarmUpBindings::cpuinfer_interface) .def("load_weights_task", &MoeBindings::LoadWeightsBindings::cpuinfer_interface) .def("forward_sft_task", &MoeBindings::ForwardSFTBindings::cpuinfer_interface) - .def("backward_task", &MoeBindings::BackwardBindings::cpuinfer_interface) + .def("backward_task", &MoeBindings::BackwardBindings::cpuinfer_interface, + py::arg("grad_output"), py::arg("grad_input"), py::arg("grad_gate_lora_a"), + py::arg("grad_gate_lora_b"), py::arg("grad_up_lora_a"), py::arg("grad_up_lora_b"), + py::arg("grad_down_lora_a"), py::arg("grad_down_lora_b"), py::arg("grad_weights"), + py::arg("grad_gate_proj"), py::arg("grad_up_proj"), py::arg("grad_down_proj"), + py::arg("accumulate_optimizer_grads") = false, py::arg("optimizer_grad_scale") = 1.0f) .def("update_lora_weights_task", &MoeBindings::UpdateLoRAWeightsBindings::cpuinfer_interface) .def("warm_up", &MoeClass::warm_up) .def("load_weights", &MoeClass::load_weights) .def("forward_sft", &MoeClass::forward_sft_binding) - .def("backward", &MoeClass::backward_binding) + .def("backward", &MoeClass::backward_binding, + py::arg("grad_output"), py::arg("grad_input"), py::arg("grad_gate_lora_a"), + py::arg("grad_gate_lora_b"), py::arg("grad_up_lora_a"), py::arg("grad_up_lora_b"), + py::arg("grad_down_lora_a"), py::arg("grad_down_lora_b"), py::arg("grad_weights"), + py::arg("grad_gate_proj"), py::arg("grad_up_proj"), py::arg("grad_down_proj"), + py::arg("accumulate_optimizer_grads") = false, py::arg("optimizer_grad_scale") = 1.0f) .def("update_lora_weights", &MoeClass::update_lora_weights_binding) .def("prepare_and_save_bwd", [](MoeClass& self, intptr_t gate, intptr_t up, intptr_t down, const std::string& path) { @@ -796,6 +812,7 @@ PYBIND11_MODULE(kt_kernel_ext, m) { .DEF_PTR_PROPERTY(MOESFTConfig, down_lora_a) .DEF_PTR_PROPERTY(MOESFTConfig, down_lora_b) .def_readwrite("full_weight_grad", &MOESFTConfig::full_weight_grad) + .def_readwrite("authoritative_optimizer_grads", &MOESFTConfig::authoritative_optimizer_grads) .DEF_PTR_PROPERTY(MOESFTConfig, grad_gate_proj) .DEF_PTR_PROPERTY(MOESFTConfig, grad_up_proj) .DEF_PTR_PROPERTY(MOESFTConfig, grad_down_proj); diff --git a/kt-kernel/operators/amx/la/bf16_dweight.hpp b/kt-kernel/operators/amx/la/bf16_dweight.hpp index 190917c59..568adc624 100644 --- a/kt-kernel/operators/amx/la/bf16_dweight.hpp +++ b/kt-kernel/operators/amx/la/bf16_dweight.hpp @@ -201,17 +201,29 @@ class BF16DWeightKernel { } static void store_bf16(const float* source, ggml_bf16_t* destination, int destination_stride, int row_count, - int column_count) { + int column_count, bool accumulate = false, float scale = 1.0f) { + const __m512 scale_vector = _mm512_set1_ps(scale); for (int row = 0; row < row_count; ++row) { const float* src_row = source + row * N_STEP; ggml_bf16_t* dst_row = destination + static_cast(row) * destination_stride; if (column_count == N_STEP) { __m512 lo = _mm512_loadu_ps(src_row); __m512 hi = _mm512_loadu_ps(src_row + 16); + lo = _mm512_mul_ps(lo, scale_vector); + hi = _mm512_mul_ps(hi, scale_vector); + if (accumulate) { + __m512 old_lo; + __m512 old_hi; + avx512_32xbf16_to_32xfp32(reinterpret_cast<__m512i*>(dst_row), &old_lo, &old_hi); + lo = _mm512_add_ps(lo, old_lo); + hi = _mm512_add_ps(hi, old_hi); + } avx512_32xfp32_to_32xbf16(&lo, &hi, reinterpret_cast<__m512i*>(dst_row)); } else { for (int column = 0; column < column_count; ++column) { - dst_row[column] = GGML_FP32_TO_BF16(src_row[column]); + float value = src_row[column] * scale; + if (accumulate) value += GGML_BF16_TO_FP32(dst_row[column]); + dst_row[column] = GGML_FP32_TO_BF16(value); } } } diff --git a/kt-kernel/operators/amx/sft_moe.hpp b/kt-kernel/operators/amx/sft_moe.hpp index c0aef51b0..e211e0956 100644 --- a/kt-kernel/operators/amx/sft_moe.hpp +++ b/kt-kernel/operators/amx/sft_moe.hpp @@ -1511,7 +1511,8 @@ class AMX_SFT_MOE_TP : public BaseMOE { void* grad_up_lora_a, void* grad_up_lora_b, void* grad_down_lora_a, void* grad_down_lora_b, void* grad_weights, int full_intermediate_size = 0, float* fp32_grad_down_lora_b = nullptr, float* fp32_grad_gate_lora_a = nullptr, float* fp32_grad_up_lora_a = nullptr, - void* grad_gate_proj = nullptr, void* grad_up_proj = nullptr, void* grad_down_proj = nullptr) { + void* grad_gate_proj = nullptr, void* grad_up_proj = nullptr, void* grad_down_proj = nullptr, + bool accumulate_optimizer_grads = false, float optimizer_grad_scale = 1.0f) { SFTProfileScope total_scope(profiler_, SFTProfileStage::BwdTotal); auto stage_start = profiler_.start(); // If full_intermediate_size not provided, use local (non-TP mode) @@ -1672,7 +1673,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { stage_start = profiler_.start(); if constexpr (supports_standard_mat_mul_v) { backward_down_amx(cache, grad_output, grad_down_lora_a, grad_down_lora_b, full_intermediate_size, - fp32_grad_down_lora_b); + fp32_grad_down_lora_b, optimizer_grad_scale); } else { // backward_down(cache, grad_output, grad_down_lora_a, grad_down_lora_b); } @@ -1797,7 +1798,8 @@ class AMX_SFT_MOE_TP : public BaseMOE { stage_start = profiler_.start(); if constexpr (supports_standard_mat_mul_v) { backward_gate_up_amx(cache, grad_input, grad_gate_lora_a, grad_gate_lora_b, grad_up_lora_a, grad_up_lora_b, - full_intermediate_size, fp32_grad_gate_lora_a, fp32_grad_up_lora_a); + full_intermediate_size, fp32_grad_gate_lora_a, fp32_grad_up_lora_a, + optimizer_grad_scale); } else { // backward_gate_up(cache, grad_input, grad_gate_lora_a, grad_gate_lora_b, grad_up_lora_a, grad_up_lora_b); } @@ -2013,7 +2015,8 @@ class AMX_SFT_MOE_TP : public BaseMOE { // ===================================================================== stage_start = profiler_.start(); if (sft_config_.full_weight_grad && grad_gate_proj && grad_up_proj && grad_down_proj) { - backward_base_weight_grad(cache, full_intermediate_size, grad_gate_proj, grad_up_proj, grad_down_proj); + backward_base_weight_grad(cache, full_intermediate_size, grad_gate_proj, grad_up_proj, grad_down_proj, + accumulate_optimizer_grads, optimizer_grad_scale); } profiler_.record(SFTProfileStage::BwdBaseWeightGrad, stage_start); @@ -2037,7 +2040,8 @@ class AMX_SFT_MOE_TP : public BaseMOE { * Uses FP32 accumulator for precision, writes BF16 output. */ void backward_base_weight_grad(const ForwardCache& cache, int full_intermediate_size, void* grad_gate_proj, - void* grad_up_proj, void* grad_down_proj) { + void* grad_up_proj, void* grad_down_proj, + bool accumulate_optimizer_grads = false, float optimizer_grad_scale = 1.0f) { auto stage_start = profiler_.start(); const int H = config_.hidden_size; const int I = config_.intermediate_size; @@ -2051,6 +2055,12 @@ class AMX_SFT_MOE_TP : public BaseMOE { auto* ggp = static_cast(grad_gate_proj); // TP slice of [E, F, H] auto* gup_ptr = static_cast(grad_up_proj); // TP slice of [E, F, H] auto* gdp = static_cast(grad_down_proj); // TP slice of [E, H, F] + auto store_optimizer_grad = [accumulate_optimizer_grads, optimizer_grad_scale](ggml_bf16_t& destination, + float current) { + float value = current * optimizer_grad_scale; + if (accumulate_optimizer_grads) value += GGML_BF16_TO_FP32(destination); + destination = GGML_FP32_TO_BF16(value); + }; std::vector expert_offsets(activated_expert); size_t token_offset = 0; @@ -2217,7 +2227,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { }); profile_operation(timings.store_ns, timings.store_calls, [&] { DWeightKernel::store_bf16(scratch.c0(), down_dst + static_cast(h_start) * F + i_start, F, - h_count, i_count); + h_count, i_count, accumulate_optimizer_grads, optimizer_grad_scale); }); } return; @@ -2241,9 +2251,9 @@ class AMX_SFT_MOE_TP : public BaseMOE { }); profile_operation(timings.store_ns, timings.store_calls, [&] { DWeightKernel::store_bf16(scratch.c0(), gate_dst + static_cast(i_start) * H + h_start, H, - i_count, h_count); + i_count, h_count, accumulate_optimizer_grads, optimizer_grad_scale); DWeightKernel::store_bf16(scratch.c1(), up_dst + static_cast(i_start) * H + h_start, H, - i_count, h_count); + i_count, h_count, accumulate_optimizer_grads, optimizer_grad_scale); }); } }, @@ -2359,7 +2369,8 @@ class AMX_SFT_MOE_TP : public BaseMOE { for (int row = 0; row < h_count; row++) { for (int col = 0; col < i_count; col++) { - down_dst[(size_t)(h_start + row) * F + i_start + col] = GGML_FP32_TO_BF16(c0[row * TILE_N + col]); + store_optimizer_grad(down_dst[(size_t)(h_start + row) * F + i_start + col], + c0[row * TILE_N + col]); } } } @@ -2422,8 +2433,10 @@ class AMX_SFT_MOE_TP : public BaseMOE { for (int row = 0; row < i_count; row++) { for (int col = 0; col < h_count; col++) { - gate_dst[(size_t)(i_start + row) * H + h_start + col] = GGML_FP32_TO_BF16(c0[row * TILE_N + col]); - up_dst[(size_t)(i_start + row) * H + h_start + col] = GGML_FP32_TO_BF16(c1[row * TILE_N + col]); + store_optimizer_grad(gate_dst[(size_t)(i_start + row) * H + h_start + col], + c0[row * TILE_N + col]); + store_optimizer_grad(up_dst[(size_t)(i_start + row) * H + h_start + col], + c1[row * TILE_N + col]); } } } @@ -2488,13 +2501,13 @@ class AMX_SFT_MOE_TP : public BaseMOE { stage_start = profiler_.start(); for (int i = 0; i < I; i++) { for (int h = 0; h < H; h++) { - ggp[(size_t)expert_idx * F * H + (size_t)i * H + h] = GGML_FP32_TO_BF16(acc_gate[i * H + h]); - gup_ptr[(size_t)expert_idx * F * H + (size_t)i * H + h] = GGML_FP32_TO_BF16(acc_up[i * H + h]); + store_optimizer_grad(ggp[(size_t)expert_idx * F * H + (size_t)i * H + h], acc_gate[i * H + h]); + store_optimizer_grad(gup_ptr[(size_t)expert_idx * F * H + (size_t)i * H + h], acc_up[i * H + h]); } } for (int h = 0; h < H; h++) { for (int i = 0; i < I; i++) { - gdp[(size_t)expert_idx * H * F + (size_t)h * F + i] = GGML_FP32_TO_BF16(acc_down[h * I + i]); + store_optimizer_grad(gdp[(size_t)expert_idx * H * F + (size_t)h * F + i], acc_down[h * I + i]); } } profiler_.record(SFTProfileStage::BwdBaseWeightGradStore, stage_start); @@ -4868,7 +4881,8 @@ class AMX_SFT_MOE_TP : public BaseMOE { */ void backward_down_amx(const ForwardCache& cache, const void* grad_output, void* grad_down_lora_a, void* grad_down_lora_b, int full_intermediate_size = 0, - float* fp32_grad_down_lora_b = nullptr) { + float* fp32_grad_down_lora_b = nullptr, + float optimizer_grad_scale = 1.0f) { auto stage_start = profiler_.start(); if (full_intermediate_size == 0) full_intermediate_size = config_.intermediate_size; auto pool = config_.pool->get_subpool(tp_part_idx); @@ -5161,7 +5175,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { max_down_lora_tokens = std::max(max_down_lora_tokens, num_tokens); } - const float scale = lora_scaling_; + const float scale = lora_scaling_ * optimizer_grad_scale; constexpr int kDownGradABBlockedThreshold = 4096; if (max_down_lora_tokens >= kDownGradABBlockedThreshold) { @@ -5636,7 +5650,8 @@ class AMX_SFT_MOE_TP : public BaseMOE { */ void backward_gate_up_amx(const ForwardCache& cache, void* grad_input, void* grad_gate_lora_a, void* grad_gate_lora_b, void* grad_up_lora_a, void* grad_up_lora_b, int full_intermediate_size = 0, - float* fp32_grad_gate_lora_a = nullptr, float* fp32_grad_up_lora_a = nullptr) { + float* fp32_grad_gate_lora_a = nullptr, float* fp32_grad_up_lora_a = nullptr, + float optimizer_grad_scale = 1.0f) { auto stage_start = profiler_.start(); if (full_intermediate_size == 0) full_intermediate_size = config_.intermediate_size; auto pool = config_.pool->get_subpool(tp_part_idx); @@ -6020,7 +6035,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { constexpr int kGuGradBBlock = 256; int gradb_blocks = (inter_size + kGuGradBBlock - 1) / kGuGradBBlock; - const float scale = lora_scaling_; + const float scale = lora_scaling_ * optimizer_grad_scale; pool->do_work_stealing_job( activated_expert * 2 * gradb_blocks, nullptr, [&, grad_gate_b, grad_up_b, activated_expert, gradb_blocks, inter_size, rank, gradb_elems, @@ -6124,7 +6139,8 @@ class AMX_SFT_MOE_TP : public BaseMOE { int grad_a_blocks = (config_.hidden_size + kGuGradATile - 1) / kGuGradATile; pool->do_work_stealing_job( activated_expert * grad_a_blocks, nullptr, - [this, do_up, grad_lora_a, fp32_grad_lora_a, use_fp32_lora_a, grad_a_blocks, &fused_bufs](int task_id) { + [this, do_up, grad_lora_a, fp32_grad_lora_a, use_fp32_lora_a, grad_a_blocks, &fused_bufs, + optimizer_grad_scale](int task_id) { int expert_task = task_id / grad_a_blocks; int block_idx = task_id % grad_a_blocks; int expert_idx = m_expert_id_map_[expert_task]; @@ -6143,7 +6159,8 @@ class AMX_SFT_MOE_TP : public BaseMOE { int tile_len = h_end - h_start; if (tile_len <= 0) return; int tile_vec_end = tile_len & ~(kVecWidth - 1); - __m512 scale_vec = _mm512_set1_ps(lora_scaling_); + const float parameter_scale = lora_scaling_ * optimizer_grad_scale; + __m512 scale_vec = _mm512_set1_ps(parameter_scale); const int lora_r = lora_rank_; // Split one expert into hidden-dimension tiles so LoRA grad_A can use all CPU threads. @@ -6189,7 +6206,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { _mm512_storeu_ps(fp32_row + h, cur); } for (; h < tile_len; h++) { - fp32_row[h] += acc_row[h] * lora_scaling_; + fp32_row[h] += acc_row[h] * parameter_scale; } } } else { @@ -6209,7 +6226,7 @@ class AMX_SFT_MOE_TP : public BaseMOE { } for (; h < tile_len; h++) { float cur = GGML_BF16_TO_FP32(grad_row[h]); - cur += acc_row[h] * lora_scaling_; + cur += acc_row[h] * parameter_scale; grad_row[h] = GGML_FP32_TO_BF16(cur); } } diff --git a/kt-kernel/operators/amx/test/test_bf16_dweight.cpp b/kt-kernel/operators/amx/test/test_bf16_dweight.cpp index 5afca4f87..8cc988965 100644 --- a/kt-kernel/operators/amx/test/test_bf16_dweight.cpp +++ b/kt-kernel/operators/amx/test/test_bf16_dweight.cpp @@ -158,6 +158,60 @@ bool run_shared_panel_case(int routes, int rows, int columns) { return passed; } +bool run_store_mode_case(int rows, int columns, int destination_stride, bool accumulate, float scale) { + constexpr int prefix_guard_elements = 5; + constexpr int suffix_guard_elements = 11; + const ggml_bf16_t guard = GGML_FP32_TO_BF16(-123.0f); + + alignas(64) float source[Kernel::M_STEP * Kernel::N_STEP]; + for (int row = 0; row < Kernel::M_STEP; ++row) { + for (int column = 0; column < Kernel::N_STEP; ++column) { + source[row * Kernel::N_STEP + column] = + static_cast((row * Kernel::N_STEP + column) % 41 - 20) * 0.03125f; + } + } + + std::vector storage( + prefix_guard_elements + static_cast(rows) * destination_stride + suffix_guard_elements, guard); + std::vector expected(static_cast(rows) * columns); + ggml_bf16_t* destination = storage.data() + prefix_guard_elements; + + for (int row = 0; row < rows; ++row) { + for (int column = 0; column < columns; ++column) { + const float initial = static_cast((row * columns + column) % 17 - 8) * 0.125f; + destination[static_cast(row) * destination_stride + column] = GGML_FP32_TO_BF16(initial); + const float old = GGML_BF16_TO_FP32(destination[static_cast(row) * destination_stride + column]); + const float value = source[row * Kernel::N_STEP + column] * scale + (accumulate ? old : 0.0f); + expected[static_cast(row) * columns + column] = GGML_FP32_TO_BF16(value); + } + } + + DWeightKernel::store_bf16(source, destination, destination_stride, rows, columns, accumulate, scale); + + bool passed = true; + for (int row = 0; row < rows; ++row) { + for (int column = 0; column < columns; ++column) { + const ggml_bf16_t actual = destination[static_cast(row) * destination_stride + column]; + if (actual.bits != expected[static_cast(row) * columns + column].bits) passed = false; + } + for (int column = columns; column < destination_stride; ++column) { + if (destination[static_cast(row) * destination_stride + column].bits != guard.bits) passed = false; + } + } + for (int i = 0; i < prefix_guard_elements; ++i) { + if (storage[i].bits != guard.bits) passed = false; + } + const size_t suffix_begin = prefix_guard_elements + static_cast(rows) * destination_stride; + for (int i = 0; i < suffix_guard_elements; ++i) { + if (storage[suffix_begin + i].bits != guard.bits) passed = false; + } + + std::printf("BF16 dWeight store mode=%s scale=%.4f shape=%dx%d stride=%d guards=%s %s\n", + accumulate ? "accumulate" : "overwrite", scale, rows, columns, destination_stride, + passed ? "intact" : "CORRUPTED", passed ? "PASS" : "FAIL"); + return passed; +} + bool run_amx_benchmark() { if constexpr (!amx::AMX_AVAILABLE) { std::printf("BF16 dWeight AMX benchmark: SKIP (AMX unavailable)\n"); @@ -309,5 +363,9 @@ int main(int argc, char** argv) { passed = run_case(65, 31, 7) && passed; passed = run_shared_panel_case(33, 45, 77) && passed; passed = run_shared_panel_case(65, 64, 96) && passed; + passed = run_store_mode_case(32, 32, 41, false, 1.0f) && passed; + passed = run_store_mode_case(32, 32, 39, true, 0.5f) && passed; + passed = run_store_mode_case(17, 29, 37, false, -0.25f) && passed; + passed = run_store_mode_case(31, 7, 19, true, 0.125f) && passed; return passed ? 0 : 1; } diff --git a/kt-kernel/operators/common.hpp b/kt-kernel/operators/common.hpp index 5c39c00d8..5c83e7b3a 100644 --- a/kt-kernel/operators/common.hpp +++ b/kt-kernel/operators/common.hpp @@ -354,6 +354,10 @@ struct MOESFTConfig : public GeneralMOEConfig { // Full weight gradient configuration bool full_weight_grad = false; + // Opt in to the C++-authoritative optimizer-gradient lifecycle. Keep this + // disabled by default so existing callers retain legacy buffer semantics. + bool authoritative_optimizer_grads = false; + // Base weight gradient buffer pointers (directly pointing to Python tensor memory, zero-copy) // Only used when full_weight_grad == true void* grad_gate_proj = nullptr; // [expert_num, intermediate_size, hidden_size] diff --git a/kt-kernel/operators/moe-sft-tp.hpp b/kt-kernel/operators/moe-sft-tp.hpp index 1468fe5b9..aae1778e5 100644 --- a/kt-kernel/operators/moe-sft-tp.hpp +++ b/kt-kernel/operators/moe-sft-tp.hpp @@ -11,6 +11,7 @@ #include #include +#include #include #include #include @@ -217,13 +218,13 @@ class TP_MOE_SFT : public TP_MOE { std::vector part_grad_input_; std::vector part_grad_weights_; - // Full-FT dWeight outputs persist across calls. Active experts are overwritten in full; - // only experts that became inactive need to be cleared after the first invocation. - std::vector last_base_grad_active_experts_; - bool base_grad_outputs_initialized_ = false; - void* last_grad_gate_proj_ = nullptr; - void* last_grad_up_proj_ = nullptr; - void* last_grad_down_proj_ = nullptr; + // The nine optimizer-gradient outputs are owned by C++ across an optimizer + // window. Python only attaches/detaches aliases to these buffers; it never + // zeros them. Keep the active-expert union from the previous completed + // window so a new window can lazily retire stale expert slices. + std::vector optimizer_grad_window_active_experts_; + std::array optimizer_grad_output_ptrs_{}; + bool optimizer_grad_outputs_initialized_ = false; public: TP_MOE_SFT(const MOESFTConfig& config) : Base(static_cast(config)), sft_config(config) { @@ -627,7 +628,8 @@ class TP_MOE_SFT : public TP_MOE { void backward(const void* grad_output, void* grad_input, void* grad_gate_lora_a, void* grad_gate_lora_b, void* grad_up_lora_a, void* grad_up_lora_b, void* grad_down_lora_a, void* grad_down_lora_b, void* grad_weights, void* grad_gate_proj = nullptr, void* grad_up_proj = nullptr, - void* grad_down_proj = nullptr) { + void* grad_down_proj = nullptr, bool accumulate_optimizer_grads = false, + float optimizer_grad_scale = 1.0f) { SFTProfileScope total_scope(profiler_, SFTProfileStage::TpBwdTotal); auto stage_start = profiler_.start(); auto pool = config.pool; @@ -643,6 +645,10 @@ class TP_MOE_SFT : public TP_MOE { const bool need_grad_weights = (grad_weights != nullptr); const bool need_base_weight_grad = sft_config.full_weight_grad && grad_gate_proj != nullptr && grad_up_proj != nullptr && grad_down_proj != nullptr; + const bool need_lora_weight_grad = !kSkipLoRA && lora_rank > 0 && grad_gate_lora_a != nullptr && + grad_gate_lora_b != nullptr && grad_up_lora_a != nullptr && + grad_up_lora_b != nullptr && grad_down_lora_a != nullptr && + grad_down_lora_b != nullptr; // SkipLoRA: zero out lora_rank to skip all LoRA buffer allocations if constexpr (kSkipLoRA) lora_rank = 0; @@ -711,7 +717,7 @@ class TP_MOE_SFT : public TP_MOE { clear_bytes[i] = offset; } - // Parallel memset for per-TP partials and final base-weight gradients. + // Parallel memset for per-TP partials and persistent optimizer gradients. struct ClearSeg { uint8_t* ptr; size_t len; @@ -732,40 +738,105 @@ class TP_MOE_SFT : public TP_MOE { } } - size_t base_grad_clear_bytes = 0; - const bool can_selectively_clear_base_grad = - need_base_weight_grad && T::kSupportsDirectBf16Reload && sft_config.lora_rank == 0 && - config.num_gpu_experts == 0; - if (need_base_weight_grad) { - const size_t base_grad_bytes = (size_t)expert_num * full_intermediate_size * hidden_size * sizeof(ggml_bf16_t); - const size_t expert_grad_bytes = (size_t)full_intermediate_size * hidden_size * sizeof(ggml_bf16_t); - auto append_clear_segments = [&](void* ptr, size_t offset, size_t bytes) { - auto* base = static_cast(ptr) + offset; - for (size_t off = 0; off < bytes; off += kChunkBytes) { - size_t len = std::min(kChunkBytes, bytes - off); - clear_segs.push_back(ClearSeg{base + off, len}); - } - }; + size_t optimizer_grad_clear_bytes = 0; + const bool supports_authoritative_optimizer_grads = + sft_config.authoritative_optimizer_grads && T::kSupportsDirectBf16Reload && config.num_gpu_experts == 0; + const bool has_authoritative_optimizer_grads = + supports_authoritative_optimizer_grads && (need_base_weight_grad || need_lora_weight_grad); + const std::array current_optimizer_grad_ptrs = { + need_base_weight_grad ? grad_gate_proj : nullptr, + need_base_weight_grad ? grad_up_proj : nullptr, + need_base_weight_grad ? grad_down_proj : nullptr, + need_lora_weight_grad ? grad_gate_lora_a : nullptr, + need_lora_weight_grad ? grad_gate_lora_b : nullptr, + need_lora_weight_grad ? grad_up_lora_a : nullptr, + need_lora_weight_grad ? grad_up_lora_b : nullptr, + need_lora_weight_grad ? grad_down_lora_a : nullptr, + need_lora_weight_grad ? grad_down_lora_b : nullptr, + }; + const bool optimizer_grad_pointers_unchanged = + current_optimizer_grad_ptrs == optimizer_grad_output_ptrs_; + const bool force_optimizer_grad_initialization = + has_authoritative_optimizer_grads && + (!optimizer_grad_outputs_initialized_ || !optimizer_grad_pointers_unchanged); + // If C++ state was invalidated by a failed invocation, clearing and + // overwriting is the only safe recovery; never accumulate into unknown data. + const bool effective_accumulate_optimizer_grads = + has_authoritative_optimizer_grads && accumulate_optimizer_grads && + !force_optimizer_grad_initialization; + + if (has_authoritative_optimizer_grads) { + profiler_.record_ns(effective_accumulate_optimizer_grads + ? SFTProfileStage::TpBwdOptimizerGradAccumulate + : SFTProfileStage::TpBwdOptimizerGradOverwrite, + 0); + // Mark invalid before dispatch. State is committed only after all NUMA + // work and TP merges complete successfully. + optimizer_grad_outputs_initialized_ = false; + } - const bool pointers_unchanged = last_grad_gate_proj_ == grad_gate_proj && last_grad_up_proj_ == grad_up_proj && - last_grad_down_proj_ == grad_down_proj; - if (!can_selectively_clear_base_grad || !base_grad_outputs_initialized_ || !pointers_unchanged) { + auto append_clear_segments = [&](void* ptr, size_t offset, size_t bytes) { + if (ptr == nullptr || bytes == 0) return; + auto* base = static_cast(ptr) + offset; + for (size_t off = 0; off < bytes; off += kChunkBytes) { + size_t len = std::min(kChunkBytes, bytes - off); + clear_segs.push_back(ClearSeg{base + off, len}); + } + optimizer_grad_clear_bytes += bytes; + }; + + std::vector current_active_mask(expert_num, 0); + for (int expert_idx : active_expert_map) { + if (expert_idx >= 0 && expert_idx < expert_num) current_active_mask[expert_idx] = 1; + } + + const size_t base_expert_bytes = + (size_t)full_intermediate_size * hidden_size * sizeof(ggml_bf16_t); + const size_t base_grad_bytes = (size_t)expert_num * base_expert_bytes; + if (need_base_weight_grad) { + if (!supports_authoritative_optimizer_grads || force_optimizer_grad_initialization) { append_clear_segments(grad_gate_proj, 0, base_grad_bytes); append_clear_segments(grad_up_proj, 0, base_grad_bytes); append_clear_segments(grad_down_proj, 0, base_grad_bytes); - base_grad_clear_bytes = 3 * base_grad_bytes; - } else { - std::vector active_mask(expert_num, 0); - for (int expert_idx : active_expert_map) { - if (expert_idx >= 0 && expert_idx < expert_num) active_mask[expert_idx] = 1; + } else if (!effective_accumulate_optimizer_grads) { + // Active experts are overwritten by dWeight. Only experts that were + // active in the previous optimizer window and are absent now are stale. + for (int expert_idx : optimizer_grad_window_active_experts_) { + if (expert_idx < 0 || expert_idx >= expert_num || current_active_mask[expert_idx]) continue; + const size_t offset = static_cast(expert_idx) * base_expert_bytes; + append_clear_segments(grad_gate_proj, offset, base_expert_bytes); + append_clear_segments(grad_up_proj, offset, base_expert_bytes); + append_clear_segments(grad_down_proj, offset, base_expert_bytes); + } + } + } + + if (need_lora_weight_grad && supports_authoritative_optimizer_grads) { + const std::array lora_ptrs = {grad_gate_lora_a, grad_gate_lora_b, grad_up_lora_a, + grad_up_lora_b, grad_down_lora_a, grad_down_lora_b}; + const std::array lora_expert_bytes = { + (size_t)lora_rank * hidden_size * sizeof(ggml_bf16_t), + (size_t)full_intermediate_size * lora_rank * sizeof(ggml_bf16_t), + (size_t)lora_rank * hidden_size * sizeof(ggml_bf16_t), + (size_t)full_intermediate_size * lora_rank * sizeof(ggml_bf16_t), + (size_t)lora_rank * full_intermediate_size * sizeof(ggml_bf16_t), + (size_t)hidden_size * lora_rank * sizeof(ggml_bf16_t), + }; + if (force_optimizer_grad_initialization) { + for (size_t output_idx = 0; output_idx < lora_ptrs.size(); ++output_idx) { + append_clear_segments(lora_ptrs[output_idx], 0, + (size_t)expert_num * lora_expert_bytes[output_idx]); } - for (int expert_idx : last_base_grad_active_experts_) { - if (expert_idx < 0 || expert_idx >= expert_num || active_mask[expert_idx]) continue; - const size_t expert_offset = static_cast(expert_idx) * expert_grad_bytes; - append_clear_segments(grad_gate_proj, expert_offset, expert_grad_bytes); - append_clear_segments(grad_up_proj, expert_offset, expert_grad_bytes); - append_clear_segments(grad_down_proj, expert_offset, expert_grad_bytes); - base_grad_clear_bytes += 3 * expert_grad_bytes; + } else if (!effective_accumulate_optimizer_grads) { + // LoRA kernels use BF16 read-modify-write, so retire every expert from + // the previous window before the first microbatch of the new window. + for (int expert_idx : optimizer_grad_window_active_experts_) { + if (expert_idx < 0 || expert_idx >= expert_num) continue; + for (size_t output_idx = 0; output_idx < lora_ptrs.size(); ++output_idx) { + append_clear_segments(lora_ptrs[output_idx], + (size_t)expert_idx * lora_expert_bytes[output_idx], + lora_expert_bytes[output_idx]); + } } } } @@ -778,35 +849,38 @@ class TP_MOE_SFT : public TP_MOE { nullptr); size_t partial_clear_bytes = 0; for (size_t bytes : clear_bytes) partial_clear_bytes += bytes; - profiler_.record_bytes(SFTProfileStage::TpBwdBufferClear, partial_clear_bytes + base_grad_clear_bytes); + profiler_.record_bytes(SFTProfileStage::TpBwdBufferClear, partial_clear_bytes + optimizer_grad_clear_bytes); + profiler_.record_bytes(SFTProfileStage::TpBwdOptimizerGradLazyClear, optimizer_grad_clear_bytes); profiler_.record(SFTProfileStage::TpBwdBufferClear, stage_start); // Compute TP-slice pointers for copy-type direct writes // Each TP writes to its own I-slice of the final output tensor - std::vector tp_gate_b_ptr(tp_count); - std::vector tp_up_b_ptr(tp_count); - std::vector tp_down_a_ptr(tp_count); - std::vector tp_fp32_down_b(tp_count); - std::vector tp_fp32_gate_a(tp_count); - std::vector tp_fp32_up_a(tp_count); + std::vector tp_gate_b_ptr(tp_count, nullptr); + std::vector tp_up_b_ptr(tp_count, nullptr); + std::vector tp_down_a_ptr(tp_count, nullptr); + std::vector tp_fp32_down_b(tp_count, nullptr); + std::vector tp_fp32_gate_a(tp_count, nullptr); + std::vector tp_fp32_up_a(tp_count, nullptr); std::vector tp_grad_gate_proj(tp_count, nullptr); std::vector tp_grad_up_proj(tp_count, nullptr); std::vector tp_grad_down_proj(tp_count, nullptr); if constexpr (!kSkipLoRA) { - int tp_offset = 0; - for (int i = 0; i < tp_count; i++) { - // Copy-type: pointer into final tensor at this TP's I-slice - tp_gate_b_ptr[i] = (ggml_bf16_t*)grad_gate_lora_b + (size_t)tp_offset * lora_rank; - tp_up_b_ptr[i] = (ggml_bf16_t*)grad_up_lora_b + (size_t)tp_offset * lora_rank; - tp_down_a_ptr[i] = (ggml_bf16_t*)grad_down_lora_a + tp_offset; // row-wise, offset added per-row + if (need_lora_weight_grad) { + int tp_offset = 0; + for (int i = 0; i < tp_count; i++) { + // Copy-type: pointer into final tensor at this TP's I-slice + tp_gate_b_ptr[i] = (ggml_bf16_t*)grad_gate_lora_b + (size_t)tp_offset * lora_rank; + tp_up_b_ptr[i] = (ggml_bf16_t*)grad_up_lora_b + (size_t)tp_offset * lora_rank; + tp_down_a_ptr[i] = (ggml_bf16_t*)grad_down_lora_a + tp_offset; // row-wise, offset added per-row - // Reduce-type: sparse FP32 partials (reinterpret from part_grad pointers) - tp_fp32_down_b[i] = (float*)part_grad_down_lora_a_[i]; // reused slot for down_lora_b FP32 - tp_fp32_gate_a[i] = (float*)part_grad_gate_lora_a_[i]; - tp_fp32_up_a[i] = (float*)part_grad_up_lora_a_[i]; + // Reduce-type: sparse FP32 partials (reinterpret from part_grad pointers) + tp_fp32_down_b[i] = (float*)part_grad_down_lora_a_[i]; // reused slot for down_lora_b FP32 + tp_fp32_gate_a[i] = (float*)part_grad_gate_lora_a_[i]; + tp_fp32_up_a[i] = (float*)part_grad_up_lora_a_[i]; - tp_offset += tp_configs[i].intermediate_size; + tp_offset += tp_configs[i].intermediate_size; + } } } @@ -823,9 +897,6 @@ class TP_MOE_SFT : public TP_MOE { throw std::runtime_error("TP intermediate_size slices do not cover the full intermediate_size"); } - // A failed dWeight computation must force a full clear on the next invocation. - if (can_selectively_clear_base_grad) base_grad_outputs_initialized_ = false; - // Run backward on each NUMA node stage_start = profiler_.start(); pool->dispense_backend()->do_numa_job([&](int numa_id) { @@ -839,16 +910,10 @@ class TP_MOE_SFT : public TP_MOE { nullptr, /* grad_down_lora_b — unused, FP32 path below */ part_grad_weights_[numa_id], full_intermediate_size, tp_fp32_down_b[numa_id], tp_fp32_gate_a[numa_id], tp_fp32_up_a[numa_id], tp_grad_gate_proj[numa_id], - tp_grad_up_proj[numa_id], tp_grad_down_proj[numa_id]); + tp_grad_up_proj[numa_id], tp_grad_down_proj[numa_id], + effective_accumulate_optimizer_grads, optimizer_grad_scale); }); profiler_.record(SFTProfileStage::TpBwdNumaCompute, stage_start); - if (can_selectively_clear_base_grad) { - last_base_grad_active_experts_ = active_expert_map; - last_grad_gate_proj_ = grad_gate_proj; - last_grad_up_proj_ = grad_up_proj; - last_grad_down_proj_ = grad_down_proj; - base_grad_outputs_initialized_ = true; - } // // Collect per-thread timing from all NUMA subpools // for (int i = 0; i < tp_count; i++) { @@ -927,76 +992,96 @@ class TP_MOE_SFT : public TP_MOE { // Copy-type grads (gate/up_lora_b, down_lora_a) were written directly — no merge needed. stage_start = profiler_.start(); if constexpr (!kSkipLoRA) { - // Sparse merge for gate_lora_a, up_lora_a: [active_count, r, H] FP32 → [E, r, H] BF16 - { - const int sparse_rows = active_count * lora_rank; // e.g. 10*8=80 vs 4096 - auto* out_gate_a = (ggml_bf16_t*)grad_gate_lora_a; - auto* out_up_a = (ggml_bf16_t*)grad_up_lora_a; - pool->do_work_stealing_job( - sparse_rows, nullptr, - [&](int sparse_row_id) { - int task = sparse_row_id / lora_rank; - int r = sparse_row_id % lora_rank; - int expert_idx = active_expert_map[task]; - size_t src_base = ((size_t)task * lora_rank + r) * hidden_size; - size_t dst_base = ((size_t)expert_idx * lora_rank + r) * hidden_size; - - ggml_bf16_t* gd = out_gate_a + dst_base; - ggml_bf16_t* ud = out_up_a + dst_base; - - int h = 0; - for (; h + 32 <= hidden_size; h += 32) { - __m512 gs0 = _mm512_loadu_ps((const float*)tp_fp32_gate_a[0] + src_base + h); - __m512 gs1 = _mm512_loadu_ps((const float*)tp_fp32_gate_a[0] + src_base + h + 16); - __m512 us0 = _mm512_loadu_ps((const float*)tp_fp32_up_a[0] + src_base + h); - __m512 us1 = _mm512_loadu_ps((const float*)tp_fp32_up_a[0] + src_base + h + 16); - for (int tp = 1; tp < tp_count; tp++) { - gs0 = _mm512_add_ps(gs0, _mm512_loadu_ps((const float*)tp_fp32_gate_a[tp] + src_base + h)); - gs1 = _mm512_add_ps(gs1, _mm512_loadu_ps((const float*)tp_fp32_gate_a[tp] + src_base + h + 16)); - us0 = _mm512_add_ps(us0, _mm512_loadu_ps((const float*)tp_fp32_up_a[tp] + src_base + h)); - us1 = _mm512_add_ps(us1, _mm512_loadu_ps((const float*)tp_fp32_up_a[tp] + src_base + h + 16)); + if (need_lora_weight_grad) { + // Sparse merge for gate_lora_a, up_lora_a: [active_count, r, H] FP32 → [E, r, H] BF16 + { + const int sparse_rows = active_count * lora_rank; // e.g. 10*8=80 vs 4096 + auto* out_gate_a = (ggml_bf16_t*)grad_gate_lora_a; + auto* out_up_a = (ggml_bf16_t*)grad_up_lora_a; + pool->do_work_stealing_job( + sparse_rows, nullptr, + [&](int sparse_row_id) { + int task = sparse_row_id / lora_rank; + int r = sparse_row_id % lora_rank; + int expert_idx = active_expert_map[task]; + size_t src_base = ((size_t)task * lora_rank + r) * hidden_size; + size_t dst_base = ((size_t)expert_idx * lora_rank + r) * hidden_size; + + ggml_bf16_t* gd = out_gate_a + dst_base; + ggml_bf16_t* ud = out_up_a + dst_base; + + int h = 0; + for (; h + 32 <= hidden_size; h += 32) { + __m512 gs0 = _mm512_loadu_ps((const float*)tp_fp32_gate_a[0] + src_base + h); + __m512 gs1 = _mm512_loadu_ps((const float*)tp_fp32_gate_a[0] + src_base + h + 16); + __m512 us0 = _mm512_loadu_ps((const float*)tp_fp32_up_a[0] + src_base + h); + __m512 us1 = _mm512_loadu_ps((const float*)tp_fp32_up_a[0] + src_base + h + 16); + for (int tp = 1; tp < tp_count; tp++) { + gs0 = _mm512_add_ps(gs0, _mm512_loadu_ps((const float*)tp_fp32_gate_a[tp] + src_base + h)); + gs1 = _mm512_add_ps(gs1, _mm512_loadu_ps((const float*)tp_fp32_gate_a[tp] + src_base + h + 16)); + us0 = _mm512_add_ps(us0, _mm512_loadu_ps((const float*)tp_fp32_up_a[tp] + src_base + h)); + us1 = _mm512_add_ps(us1, _mm512_loadu_ps((const float*)tp_fp32_up_a[tp] + src_base + h + 16)); + } + // Authoritative BF16 buffers persist across microbatches. Legacy + // backends keep their original reduce-merge overwrite semantics. + if (supports_authoritative_optimizer_grads) { + __m512 old_gs0, old_gs1, old_us0, old_us1; + avx512_32xbf16_to_32xfp32((__m512i*)(gd + h), &old_gs0, &old_gs1); + avx512_32xbf16_to_32xfp32((__m512i*)(ud + h), &old_us0, &old_us1); + gs0 = _mm512_add_ps(gs0, old_gs0); + gs1 = _mm512_add_ps(gs1, old_gs1); + us0 = _mm512_add_ps(us0, old_us0); + us1 = _mm512_add_ps(us1, old_us1); + } + avx512_32xfp32_to_32xbf16(&gs0, &gs1, (__m512i*)(gd + h)); + avx512_32xfp32_to_32xbf16(&us0, &us1, (__m512i*)(ud + h)); } - avx512_32xfp32_to_32xbf16(&gs0, &gs1, (__m512i*)(gd + h)); - avx512_32xfp32_to_32xbf16(&us0, &us1, (__m512i*)(ud + h)); - } - for (; h < hidden_size; h++) { - float gs = ((const float*)tp_fp32_gate_a[0])[src_base + h]; - float us = ((const float*)tp_fp32_up_a[0])[src_base + h]; - for (int tp = 1; tp < tp_count; tp++) { - gs += ((const float*)tp_fp32_gate_a[tp])[src_base + h]; - us += ((const float*)tp_fp32_up_a[tp])[src_base + h]; + for (; h < hidden_size; h++) { + float gs = ((const float*)tp_fp32_gate_a[0])[src_base + h]; + float us = ((const float*)tp_fp32_up_a[0])[src_base + h]; + for (int tp = 1; tp < tp_count; tp++) { + gs += ((const float*)tp_fp32_gate_a[tp])[src_base + h]; + us += ((const float*)tp_fp32_up_a[tp])[src_base + h]; + } + if (supports_authoritative_optimizer_grads) { + gs += GGML_BF16_TO_FP32(gd[h]); + us += GGML_BF16_TO_FP32(ud[h]); + } + gd[h] = GGML_FP32_TO_BF16(gs); + ud[h] = GGML_FP32_TO_BF16(us); } - gd[h] = GGML_FP32_TO_BF16(gs); - ud[h] = GGML_FP32_TO_BF16(us); - } - }, - nullptr); - } + }, + nullptr); + } - // Sparse merge for down_lora_b: [active_count, H, r] FP32 → [E, H, r] BF16 - { - const int sparse_rows = active_count; // one task per active expert - auto* out_down_b = (ggml_bf16_t*)grad_down_lora_b; - pool->do_work_stealing_job( - sparse_rows, nullptr, - [&](int task) { - int expert_idx = active_expert_map[task]; - size_t src_expert_base = (size_t)task * hidden_size * lora_rank; - size_t dst_expert_base = (size_t)expert_idx * hidden_size * lora_rank; - - for (int hh = 0; hh < hidden_size; hh++) { - size_t src_row = src_expert_base + (size_t)hh * lora_rank; - size_t dst_row = dst_expert_base + (size_t)hh * lora_rank; - for (int r = 0; r < lora_rank; r++) { - float sum = ((const float*)tp_fp32_down_b[0])[src_row + r]; - for (int tp = 1; tp < tp_count; tp++) { - sum += ((const float*)tp_fp32_down_b[tp])[src_row + r]; + // Sparse merge for down_lora_b: [active_count, H, r] FP32 → [E, H, r] BF16 + { + const int sparse_rows = active_count; // one task per active expert + auto* out_down_b = (ggml_bf16_t*)grad_down_lora_b; + pool->do_work_stealing_job( + sparse_rows, nullptr, + [&](int task) { + int expert_idx = active_expert_map[task]; + size_t src_expert_base = (size_t)task * hidden_size * lora_rank; + size_t dst_expert_base = (size_t)expert_idx * hidden_size * lora_rank; + + for (int hh = 0; hh < hidden_size; hh++) { + size_t src_row = src_expert_base + (size_t)hh * lora_rank; + size_t dst_row = dst_expert_base + (size_t)hh * lora_rank; + for (int r = 0; r < lora_rank; r++) { + float sum = ((const float*)tp_fp32_down_b[0])[src_row + r]; + for (int tp = 1; tp < tp_count; tp++) { + sum += ((const float*)tp_fp32_down_b[tp])[src_row + r]; + } + if (supports_authoritative_optimizer_grads) { + sum += GGML_BF16_TO_FP32(out_down_b[dst_row + r]); + } + out_down_b[dst_row + r] = GGML_FP32_TO_BF16(sum); } - out_down_b[dst_row + r] = GGML_FP32_TO_BF16(sum); } - } - }, - nullptr); + }, + nullptr); + } } } // if constexpr (!kSkipLoRA) profiler_.record(SFTProfileStage::TpBwdLoraMerge, stage_start); @@ -1042,6 +1127,31 @@ class TP_MOE_SFT : public TP_MOE { profiler_.record(SFTProfileStage::TpBwdRouterGradMerge, stage_start); pool->dispense_backend()->do_numa_job([&](int numa_id) {}); + + if (has_authoritative_optimizer_grads) { + if (effective_accumulate_optimizer_grads) { + std::vector window_mask(expert_num, 0); + for (int expert_idx : optimizer_grad_window_active_experts_) { + if (expert_idx >= 0 && expert_idx < expert_num) window_mask[expert_idx] = 1; + } + for (int expert_idx : active_expert_map) { + if (expert_idx >= 0 && expert_idx < expert_num) window_mask[expert_idx] = 1; + } + optimizer_grad_window_active_experts_.clear(); + for (int expert_idx = 0; expert_idx < expert_num; ++expert_idx) { + if (window_mask[expert_idx]) optimizer_grad_window_active_experts_.push_back(expert_idx); + } + } else { + optimizer_grad_window_active_experts_ = active_expert_map; + std::sort(optimizer_grad_window_active_experts_.begin(), optimizer_grad_window_active_experts_.end()); + optimizer_grad_window_active_experts_.erase( + std::unique(optimizer_grad_window_active_experts_.begin(), + optimizer_grad_window_active_experts_.end()), + optimizer_grad_window_active_experts_.end()); + } + optimizer_grad_output_ptrs_ = current_optimizer_grad_ptrs; + optimizer_grad_outputs_initialized_ = true; + } } /** @@ -1050,10 +1160,12 @@ class TP_MOE_SFT : public TP_MOE { void backward_binding(intptr_t grad_output, intptr_t grad_input, intptr_t grad_gate_lora_a, intptr_t grad_gate_lora_b, intptr_t grad_up_lora_a, intptr_t grad_up_lora_b, intptr_t grad_down_lora_a, intptr_t grad_down_lora_b, intptr_t grad_weights, intptr_t grad_gate_proj, - intptr_t grad_up_proj, intptr_t grad_down_proj) { + intptr_t grad_up_proj, intptr_t grad_down_proj, + bool accumulate_optimizer_grads = false, float optimizer_grad_scale = 1.0f) { backward((const void*)grad_output, (void*)grad_input, (void*)grad_gate_lora_a, (void*)grad_gate_lora_b, (void*)grad_up_lora_a, (void*)grad_up_lora_b, (void*)grad_down_lora_a, (void*)grad_down_lora_b, - (void*)grad_weights, (void*)grad_gate_proj, (void*)grad_up_proj, (void*)grad_down_proj); + (void*)grad_weights, (void*)grad_gate_proj, (void*)grad_up_proj, (void*)grad_down_proj, + accumulate_optimizer_grads, optimizer_grad_scale); } /** diff --git a/kt-kernel/operators/sft_profile.hpp b/kt-kernel/operators/sft_profile.hpp index 22dd6b1cf..0785aa71b 100644 --- a/kt-kernel/operators/sft_profile.hpp +++ b/kt-kernel/operators/sft_profile.hpp @@ -72,6 +72,9 @@ enum class SFTProfileStage : uint8_t { TpFwdNumaCompute, TpFwdMerge, TpBwdTotal, + TpBwdOptimizerGradOverwrite, + TpBwdOptimizerGradAccumulate, + TpBwdOptimizerGradLazyClear, TpBwdBufferClear, TpBwdNumaCompute, TpBwdGradInputMerge, @@ -143,6 +146,9 @@ inline constexpr std::array(SFTProfileStage::Co "tp.forward.numa_compute", "tp.forward.merge", "tp.backward.total", + "tp.backward.optimizer_grad.overwrite", + "tp.backward.optimizer_grad.accumulate", + "tp.backward.optimizer_grad.lazy_clear", "tp.backward.buffer_clear", "tp.backward.numa_compute", "tp.backward.grad_input_merge", diff --git a/kt-kernel/python/sft/amx.py b/kt-kernel/python/sft/amx.py index c59a8b455..bb6710a64 100644 --- a/kt-kernel/python/sft/amx.py +++ b/kt-kernel/python/sft/amx.py @@ -41,7 +41,7 @@ AMXInt8_SFT_MOE_SkipLoRA = None AMXInt4_SFT_MOE_SkipLoRA = None -from .base import BaseSFTMoEWrapper, KExpertsSFTBuffer +from .base import BaseSFTMoEWrapper, KExpertsSFTBuffer, _supports_authoritative_optimizer_grads _AMX_M_STEP = 32 @@ -111,6 +111,12 @@ def __init__( self.method = method self._is_skip_lora = "SkipLoRA" in method + # Authoritative optimizer gradients currently rely on the BF16 SFT + # kernel's overwrite/accumulate/lazy-clear implementation. Quantized + # and SkipLoRA backends intentionally retain their legacy lifecycle. + self._uses_authoritative_optimizer_grads = _supports_authoritative_optimizer_grads( + method, self.num_gpu_experts + ) self.group_size = group_size self.zero_point = zero_point @@ -140,7 +146,12 @@ def _make_forward_task(self, buffer: KExpertsSFTBuffer, save_for_backward: bool) save_for_backward, ) - def _make_backward_task(self, buffer: KExpertsSFTBuffer): + def _make_backward_task( + self, + buffer: KExpertsSFTBuffer, + accumulate_optimizer_grads: bool = False, + optimizer_grad_scale: float = 1.0, + ): if self._is_skip_lora: return self.moe.backward_task( buffer.grad_output_cpu.data_ptr(), @@ -168,7 +179,7 @@ def _make_backward_task(self, buffer: KExpertsSFTBuffer): self.grad_down_proj_buf.data_ptr() if self._full_weight_grad and self.grad_down_proj_buf is not None else 0 ) - return self.moe.backward_task( + backward_args = ( buffer.grad_output_cpu.data_ptr(), buffer.grad_input_cpu.data_ptr(), self.grad_gate_lora_a.data_ptr() if self.lora_rank > 0 else 0, @@ -182,6 +193,13 @@ def _make_backward_task(self, buffer: KExpertsSFTBuffer): grad_up_proj_ptr, grad_down_proj_ptr, ) + if self._uses_authoritative_optimizer_grads: + return self.moe.backward_task( + *backward_args, + bool(accumulate_optimizer_grads), + float(optimizer_grad_scale), + ) + return self.moe.backward_task(*backward_args) # ========== Weight loading ========== @@ -216,6 +234,7 @@ def load_weights(self, physical_to_logical_map_cpu: torch.Tensor) -> None: config.share_backward_bb = getattr(self, "share_backward_bb", False) config.share_cache_pool = getattr(self, "share_cache_pool", False) config.full_weight_grad = self._full_weight_grad + config.authoritative_optimizer_grads = self._uses_authoritative_optimizer_grads config.physical_to_logical_map = self._physical_to_logical_map_cpu.data_ptr() if getattr(self, "_use_kt_direct_load", False): diff --git a/kt-kernel/python/sft/autograd.py b/kt-kernel/python/sft/autograd.py index 615429a7c..f3ea41fd4 100644 --- a/kt-kernel/python/sft/autograd.py +++ b/kt-kernel/python/sft/autograd.py @@ -146,6 +146,9 @@ def forward( ctx.full_weight_grad = ( wrapper is not None and getattr(wrapper, "_full_weight_grad", False) and gate_proj_param is not None ) + ctx.authoritative_optimizer_grads = bool( + wrapper is not None and getattr(wrapper, "_uses_authoritative_optimizer_grads", False) + ) # Save a sentinel tensor so non-reentrant checkpoint's saved_tensors # hooks can intercept it. When backward accesses ctx.saved_tensors, @@ -211,72 +214,123 @@ def backward(ctx, grad_output: torch.Tensor): rank=rank, world_size=world_size, ) + authoritative_grad_published = False + + def close_published_authoritative_window() -> None: + if not (rank == 0 and authoritative_grad_published and ctx.wrapper is not None): + return + try: + ctx.wrapper.release_authoritative_optimizer_grads() + except Exception: + logger.exception( + "Failed to close authoritative optimizer-gradient window after distributed backward error" + ) + if rank == 0: all_go = torch.cat(gathered_go, dim=0) total_qlen = int(all_go.shape[0]) with torch.profiler.record_function("kt.sft.cpu_backward"): - backward_out = ctx.wrapper.backward( - all_go, - output_device=ctx.original_device, + if ctx.authoritative_optimizer_grads: + backward_out = ctx.wrapper.backward( + all_go, + output_device=ctx.original_device, + optimizer_grad_scale=1.0 / world_size, + ) + else: + backward_out = ctx.wrapper.backward( + all_go, + output_device=ctx.original_device, + ) + authoritative_grad_published = ctx.authoritative_optimizer_grads + try: + if isinstance(backward_out, tuple) and len(backward_out) == 2: + all_grad_input, all_grad_weights = backward_out + elif isinstance(backward_out, tuple) and len(backward_out) == 3: + all_grad_input, _, all_grad_weights = backward_out + else: + raise ValueError("KTMoEWrapper.backward returned unexpected format.") + + all_grad_input = all_grad_input.to(dtype=ctx.original_dtype).view(total_qlen, hidden_size) + all_grad_weights = all_grad_weights.to(dtype=torch.bfloat16).view( + total_qlen, num_experts_per_tok ) - if isinstance(backward_out, tuple) and len(backward_out) == 2: - all_grad_input, all_grad_weights = backward_out - elif isinstance(backward_out, tuple) and len(backward_out) == 3: - all_grad_input, _, all_grad_weights = backward_out - else: - raise ValueError("KTMoEWrapper.backward returned unexpected format.") - - all_grad_input = all_grad_input.to(dtype=ctx.original_dtype).view(total_qlen, hidden_size) - all_grad_weights = all_grad_weights.to(dtype=torch.bfloat16).view(total_qlen, num_experts_per_tok) - offsets = _qlen_offsets(all_qlens) - scatter_gi = [all_grad_input[offsets[i] : offsets[i + 1]].contiguous() for i in range(world_size)] - scatter_gw = [all_grad_weights[offsets[i] : offsets[i + 1]].contiguous() for i in range(world_size)] + offsets = _qlen_offsets(all_qlens) + scatter_gi = [ + all_grad_input[offsets[i] : offsets[i + 1]].contiguous() for i in range(world_size) + ] + scatter_gw = [ + all_grad_weights[offsets[i] : offsets[i + 1]].contiguous() for i in range(world_size) + ] + except Exception: + close_published_authoritative_window() + raise else: scatter_gi = None scatter_gw = None - grad_input_flat = _dist_scatter_varlen_from_rank0( - rank0_chunks=scatter_gi, - all_qlens=all_qlens, - rank=rank, - world_size=world_size, - feature_shape=(hidden_size,), - device=ctx.original_device, - dtype=ctx.original_dtype, - ) - grad_weights_flat = _dist_scatter_varlen_from_rank0( - rank0_chunks=scatter_gw, - all_qlens=all_qlens, - rank=rank, - world_size=world_size, - feature_shape=(num_experts_per_tok,), - device=ctx.weights_device, - dtype=torch.bfloat16, - ) - grad_input = grad_input_flat.view(batch_size, seq_len, hidden_size) - grad_weights = grad_weights_flat.view(ctx.weights_shape).to(dtype=ctx.weights_dtype) + try: + grad_input_flat = _dist_scatter_varlen_from_rank0( + rank0_chunks=scatter_gi, + all_qlens=all_qlens, + rank=rank, + world_size=world_size, + feature_shape=(hidden_size,), + device=ctx.original_device, + dtype=ctx.original_dtype, + ) + grad_weights_flat = _dist_scatter_varlen_from_rank0( + rank0_chunks=scatter_gw, + all_qlens=all_qlens, + rank=rank, + world_size=world_size, + feature_shape=(num_experts_per_tok,), + device=ctx.weights_device, + dtype=torch.bfloat16, + ) + grad_input = grad_input_flat.view(batch_size, seq_len, hidden_size) + grad_weights = grad_weights_flat.view(ctx.weights_shape).to(dtype=ctx.weights_dtype) + except Exception: + close_published_authoritative_window() + raise elif not ctx.use_broadcast: # ---- Single-GPU path ---- grad_output_flat = grad_output.view(qlen, hidden_size) try: - with torch.profiler.record_function("kt.sft.cpu_backward"): - backward_out = ctx.wrapper.backward( - grad_output_flat, - output_device=ctx.original_device, - ) - finally: - ctx.wrapper.clear_checkpoint_output() - if isinstance(backward_out, tuple) and len(backward_out) == 2: - grad_input, grad_weights = backward_out - elif isinstance(backward_out, tuple) and len(backward_out) == 3: - grad_input, _, grad_weights = backward_out - else: - raise ValueError("KTMoEWrapper.backward returned unexpected format.") - grad_input = grad_input.view(batch_size, seq_len, hidden_size).to(dtype=ctx.original_dtype) - grad_weights = grad_weights.to(dtype=torch.bfloat16) + try: + with torch.profiler.record_function("kt.sft.cpu_backward"): + if ctx.authoritative_optimizer_grads: + backward_out = ctx.wrapper.backward( + grad_output_flat, + output_device=ctx.original_device, + optimizer_grad_scale=1.0, + ) + else: + backward_out = ctx.wrapper.backward( + grad_output_flat, + output_device=ctx.original_device, + ) + finally: + ctx.wrapper.clear_checkpoint_output() + if isinstance(backward_out, tuple) and len(backward_out) == 2: + grad_input, grad_weights = backward_out + elif isinstance(backward_out, tuple) and len(backward_out) == 3: + grad_input, _, grad_weights = backward_out + else: + raise ValueError("KTMoEWrapper.backward returned unexpected format.") + grad_input = grad_input.view(batch_size, seq_len, hidden_size).to(dtype=ctx.original_dtype) + grad_weights = grad_weights.to(dtype=torch.bfloat16) + except Exception: + if ctx.authoritative_optimizer_grads and ctx.wrapper is not None: + try: + ctx.wrapper.release_authoritative_optimizer_grads() + except Exception: + logger.exception( + "Failed to close authoritative optimizer-gradient window after local backward error" + ) + raise else: # No wrapper, no dist — shouldn't happen in normal flow grad_input = torch.zeros( @@ -287,11 +341,23 @@ def backward(ctx, grad_output: torch.Tensor): # Trigger async repack for next MoE layer in backward order next_bwd = getattr(ctx.wrapper, "_next_backward_wrapper", None) if next_bwd is not None and getattr(next_bwd, "share_backward_bb", False): - with torch.profiler.record_function("kt.sft.submit_backward_repack"): - next_bwd.submit_backward_repack() - - # Base weight gradients: return C++-written grad buffers in full mode, None otherwise - if ctx.full_weight_grad and ctx.wrapper is not None: + try: + with torch.profiler.record_function("kt.sft.submit_backward_repack"): + next_bwd.submit_backward_repack() + except Exception: + if ctx.authoritative_optimizer_grads and ctx.wrapper is not None: + try: + ctx.wrapper.release_authoritative_optimizer_grads() + except Exception: + logger.exception( + "Failed to close authoritative optimizer-gradient window after repack submission error" + ) + raise + + # Legacy backends still use PyTorch AccumulateGrad. AMXBF16_SFT + # publishes the C++ buffers directly as Parameter.grad, so returning + # them here would create a second giant copy and an aten::add_. + if ctx.full_weight_grad and ctx.wrapper is not None and not ctx.authoritative_optimizer_grads: grad_gate_proj = ctx.wrapper.grad_gate_proj_buf grad_up_proj = ctx.wrapper.grad_up_proj_buf grad_down_proj = ctx.wrapper.grad_down_proj_buf diff --git a/kt-kernel/python/sft/base.py b/kt-kernel/python/sft/base.py index d402dc3fe..9d14fd331 100644 --- a/kt-kernel/python/sft/base.py +++ b/kt-kernel/python/sft/base.py @@ -11,6 +11,8 @@ from __future__ import annotations +from dataclasses import dataclass +import math import torch from typing import Optional, Tuple from abc import ABC, abstractmethod @@ -18,6 +20,35 @@ from ..experts_base import KExpertsCPUBuffer, _MoEBase +def _supports_authoritative_optimizer_grads(method: str, num_gpu_experts: int) -> bool: + """Whether this backend can use C++-authoritative optimizer gradients.""" + return method == "AMXBF16_SFT" and int(num_gpu_experts) == 0 + + +@dataclass(frozen=True) +class _AuthoritativeOptimizerGrad: + """Stable Parameter-to-C++-gradient binding for one optimizer tensor.""" + + name: str + parameter: torch.nn.Parameter + grad_view: torch.Tensor + metadata: tuple + + +def _authoritative_grad_metadata(tensor: torch.Tensor) -> tuple: + device_index = tensor.device.index if tensor.device.index is not None else -1 + return ( + tensor.dtype, + tensor.layout, + tensor.device.type, + device_index, + tuple(tensor.shape), + tuple(tensor.stride()), + int(tensor.storage_offset()), + int(tensor.data_ptr()), + ) + + class KExpertsSFTBuffer: """ CPU buffer management for SFT expert computation. @@ -194,6 +225,10 @@ def __init__( self._cache_depth: int = 0 self._is_skip_lora: bool = False self._base_weights_dirty: bool = False + # AMXSFTMoEWrapper enables this capability only for AMXBF16_SFT. + # Keeping it false here preserves legacy INT8/INT4/SkipLoRA behavior. + self._uses_authoritative_optimizer_grads: bool = False + self._init_authoritative_optimizer_grads() self.reuse_checkpoint_forward: bool = False self._kt_has_cached_forward: bool = False self._checkpoint_output_cpu: Optional[torch.Tensor] = None @@ -201,6 +236,152 @@ def __init__( self.moe = None + # ========== Authoritative optimizer-gradient lifecycle ========== + + def _init_authoritative_optimizer_grads(self) -> None: + self._authoritative_optimizer_grads: list[_AuthoritativeOptimizerGrad] = [] + self._authoritative_grad_submission_pending: bool = False + self._authoritative_grad_pending_accumulate: bool = False + + @property + def authoritative_optimizer_grads(self) -> tuple[_AuthoritativeOptimizerGrad, ...]: + """Read-only view of the persistent C++ optimizer-gradient bindings.""" + return tuple(self._authoritative_optimizer_grads) + + def register_authoritative_optimizer_grad( + self, + name: str, + parameter: torch.nn.Parameter, + grad_view: torch.Tensor, + ) -> None: + """Register a C++-written gradient as the Parameter's sole grad copy.""" + if not self._uses_authoritative_optimizer_grads: + return + if not isinstance(parameter, torch.nn.Parameter): + raise TypeError(f"{name}: expected nn.Parameter, got {type(parameter)!r}") + if not isinstance(grad_view, torch.Tensor): + raise TypeError(f"{name}: expected Tensor grad view, got {type(grad_view)!r}") + if parameter.shape != grad_view.shape: + raise ValueError( + f"{name}: Parameter/grad shape mismatch: {tuple(parameter.shape)} != {tuple(grad_view.shape)}" + ) + if parameter.dtype != grad_view.dtype or parameter.device != grad_view.device: + raise ValueError( + f"{name}: Parameter/grad dtype or device mismatch: " + f"{parameter.dtype}/{parameter.device} != {grad_view.dtype}/{grad_view.device}" + ) + for entry in self._authoritative_optimizer_grads: + if entry.parameter is parameter: + raise RuntimeError(f"{name}: Parameter is already registered as {entry.name}") + if entry.grad_view is grad_view: + raise RuntimeError(f"{name}: grad view is already registered as {entry.name}") + + # Initial state is always a closed optimizer window. The view remains + # alive in this registry while Parameter.grad is None. + parameter.grad = None + self._authoritative_optimizer_grads.append( + _AuthoritativeOptimizerGrad( + name=name, + parameter=parameter, + grad_view=grad_view, + metadata=_authoritative_grad_metadata(grad_view), + ) + ) + + def _validate_authoritative_grad_metadata(self) -> None: + for entry in self._authoritative_optimizer_grads: + current = _authoritative_grad_metadata(entry.grad_view) + if current != entry.metadata: + raise RuntimeError( + f"{entry.name}: authoritative grad view metadata changed; " + "replacing or resizing a C++ gradient buffer is unsupported" + ) + parameter = entry.parameter + grad_view = entry.grad_view + if parameter.shape != grad_view.shape or parameter.dtype != grad_view.dtype: + raise RuntimeError(f"{entry.name}: Parameter metadata no longer matches its authoritative grad view") + if parameter.device != grad_view.device: + raise RuntimeError(f"{entry.name}: Parameter device no longer matches its authoritative grad view") + + def validate_authoritative_optimizer_grad_state(self) -> str: + """Validate aliases and return ``empty``, ``closed``, or ``open``.""" + if not self._uses_authoritative_optimizer_grads or not self._authoritative_optimizer_grads: + return "empty" + self._validate_authoritative_grad_metadata() + + none_count = 0 + alias_count = 0 + for entry in self._authoritative_optimizer_grads: + grad = entry.parameter.grad + if grad is None: + none_count += 1 + elif grad is entry.grad_view: + alias_count += 1 + else: + raise RuntimeError( + f"{entry.name}: Parameter.grad was externally replaced; " + "expected None or the registered authoritative grad view" + ) + + total = len(self._authoritative_optimizer_grads) + if none_count == total: + return "closed" + if alias_count == total: + return "open" + raise RuntimeError( + "Mixed authoritative optimizer-gradient state: some KT Parameter.grad values are None " + "while others still alias their C++ buffers" + ) + + def _prepare_authoritative_optimizer_grad_write(self, optimizer_grad_scale: float) -> bool: + if self._authoritative_grad_submission_pending: + raise RuntimeError("An authoritative optimizer-gradient backward submission is already pending") + scale = float(optimizer_grad_scale) + if not math.isfinite(scale) or scale <= 0.0: + raise ValueError(f"optimizer_grad_scale must be finite and positive, got {optimizer_grad_scale}") + state = self.validate_authoritative_optimizer_grad_state() + accumulate = state == "open" + self._authoritative_grad_submission_pending = True + self._authoritative_grad_pending_accumulate = accumulate + return accumulate + + def _publish_authoritative_optimizer_grads(self) -> None: + if not self._authoritative_grad_submission_pending: + raise RuntimeError("No authoritative optimizer-gradient backward submission is pending") + expected_state = "open" if self._authoritative_grad_pending_accumulate else "closed" + try: + state = self.validate_authoritative_optimizer_grad_state() + if state not in ("empty", expected_state): + raise RuntimeError( + f"Authoritative optimizer-gradient state changed during C++ backward: " + f"expected {expected_state}, got {state}" + ) + for entry in self._authoritative_optimizer_grads: + entry.parameter.grad = entry.grad_view + except Exception: + self._abort_authoritative_optimizer_grad_write() + raise + self._authoritative_grad_submission_pending = False + self._authoritative_grad_pending_accumulate = False + + def _abort_authoritative_optimizer_grad_write(self) -> None: + # A failed C++ task may have partially modified its outputs. Closing + # the Python window forces the next task down the overwrite/full-init path. + for entry in self._authoritative_optimizer_grads: + entry.parameter.grad = None + self._authoritative_grad_submission_pending = False + self._authoritative_grad_pending_accumulate = False + + def release_authoritative_optimizer_grads(self) -> None: + """Close the optimizer window after step without touching C++ buffers.""" + if not self._uses_authoritative_optimizer_grads: + return + if self._authoritative_grad_submission_pending: + raise RuntimeError("Cannot release authoritative gradients while C++ backward is pending") + self.validate_authoritative_optimizer_grad_state() + for entry in self._authoritative_optimizer_grads: + entry.parameter.grad = None + @staticmethod def _validate_sft_config( lora_rank: int, lora_alpha: float, max_cache_depth: int, full_weight_grad: bool = False @@ -244,10 +425,12 @@ def init_full_weight_grad_buffers( self.grad_up_proj_buf = torch.zeros(E, I, H, dtype=dtype, device="cpu") self.grad_down_proj_buf = torch.zeros(E, H, I, dtype=dtype, device="cpu") - # Note: .grad is NOT pre-assigned here. PyTorch autograd will set it - # when KTMoEFunction.backward() returns the gradient buffers. - # The C++ kernel writes directly to grad_gate_proj_buf etc., - # and backward returns them so PyTorch can propagate correctly. + if self._uses_authoritative_optimizer_grads: + self.register_authoritative_optimizer_grad("base.gate_proj", self.gate_proj_buf, self.grad_gate_proj_buf) + self.register_authoritative_optimizer_grad("base.up_proj", self.up_proj_buf, self.grad_up_proj_buf) + self.register_authoritative_optimizer_grad("base.down_proj", self.down_proj_buf, self.grad_down_proj_buf) + # Legacy backends leave .grad unset here and return these buffers from + # KTMoEFunction.backward(), preserving their existing AccumulateGrad path. @abstractmethod def update_base_weights(self) -> None: @@ -262,7 +445,12 @@ def _make_forward_task(self, buffer: KExpertsSFTBuffer, save_for_backward: bool) ... @abstractmethod - def _make_backward_task(self, buffer: KExpertsSFTBuffer): + def _make_backward_task( + self, + buffer: KExpertsSFTBuffer, + accumulate_optimizer_grads: bool = False, + optimizer_grad_scale: float = 1.0, + ): """Construct the C++ backward task object. Backend-specific.""" ... @@ -305,7 +493,9 @@ def _get_buffer(self, qlen: int) -> KExpertsSFTBuffer: def _validate_forward_inputs(self, hidden_states: torch.Tensor, expert_ids: torch.Tensor, weights: torch.Tensor): if not self._weights_loaded: raise RuntimeError("Weights not loaded. Call load_weights() or load_weights_from_tensors() first.") - if not self._lora_initialized and not self._is_skip_lora and not self._full_weight_grad: + # Hybrid mode still requires LoRA buffers even though base gradients are + # enabled. Only pure Full (lora_rank == 0) may legitimately skip them. + if self.lora_rank > 0 and not self._lora_initialized and not self._is_skip_lora: raise RuntimeError("LoRA weights not initialized. Call init_lora_weights() first.") qlen = hidden_states.shape[0] if qlen > self.chunked_prefill_size: @@ -412,20 +602,47 @@ def backward( self, grad_output: torch.Tensor, output_device: Optional[torch.device] = None, + optimizer_grad_scale: float = 1.0, ) -> Tuple[torch.Tensor, torch.Tensor]: """Backward pass computing grad_input and grad_weights.""" if self._cache_depth <= 0: raise RuntimeError("No forward cache available. Call forward(save_for_backward=True) first.") + if self._uses_authoritative_optimizer_grads and self._authoritative_grad_submission_pending: + raise RuntimeError("An authoritative optimizer-gradient backward submission is already pending") qlen = grad_output.shape[0] buffer = self._get_buffer(qlen) self._copy_grad_output_to_cpu(buffer, grad_output, qlen) - self.cpu_infer.submit(self._make_backward_task(buffer)) - self.cpu_infer.sync() + use_authoritative = self._uses_authoritative_optimizer_grads + accumulate_optimizer_grads = False + if use_authoritative: + accumulate_optimizer_grads = self._prepare_authoritative_optimizer_grad_write(optimizer_grad_scale) + try: + if use_authoritative: + backward_task = self._make_backward_task( + buffer, + accumulate_optimizer_grads=accumulate_optimizer_grads, + optimizer_grad_scale=optimizer_grad_scale, + ) + else: + backward_task = self._make_backward_task(buffer) + self.cpu_infer.submit(backward_task) + self.cpu_infer.sync() + result = self._return_grads(buffer, qlen, output_device) + if use_authoritative: + self._publish_authoritative_optimizer_grads() + except Exception: + if use_authoritative: + self._abort_authoritative_optimizer_grad_write() + # The C++ forward cache may already have been consumed by a + # partially executed backward. Require a fresh forward before + # retrying instead of reusing an indeterminate cache entry. + self._cache_depth = max(0, self._cache_depth - 1) + raise self._cache_depth -= 1 - return self._return_grads(buffer, qlen, output_device) + return result # ========== Async forward ========== @@ -572,29 +789,78 @@ def submit_backward_async( self, grad_output: torch.Tensor, output_device: Optional[torch.device] = None, + optimizer_grad_scale: float = 1.0, ) -> None: """Submit backward task without waiting. Call sync_backward() for results.""" if self._cache_depth <= 0: raise RuntimeError("No forward cache available. Call forward(save_for_backward=True) first.") + if self._uses_authoritative_optimizer_grads and self._authoritative_grad_submission_pending: + raise RuntimeError("An authoritative optimizer-gradient backward submission is already pending") qlen = grad_output.shape[0] buffer = self._get_buffer(qlen) self._copy_grad_output_to_cpu(buffer, grad_output, qlen) - self.cpu_infer.submit(self._make_backward_task(buffer)) + use_authoritative = self._uses_authoritative_optimizer_grads + accumulate_optimizer_grads = False + if use_authoritative: + accumulate_optimizer_grads = self._prepare_authoritative_optimizer_grad_write(optimizer_grad_scale) + try: + if use_authoritative: + backward_task = self._make_backward_task( + buffer, + accumulate_optimizer_grads=accumulate_optimizer_grads, + optimizer_grad_scale=optimizer_grad_scale, + ) + else: + backward_task = self._make_backward_task(buffer) + self.cpu_infer.submit(backward_task) + except Exception: + if use_authoritative: + # submit() is expected to be atomic, but drain defensively in + # case a backend queued work before reporting an error. + try: + self.cpu_infer.sync() + except Exception: + pass + self._abort_authoritative_optimizer_grad_write() + self._cache_depth = max(0, self._cache_depth - 1) + self._async_bwd_qlen = None + self._async_bwd_output_device = None + self._async_bwd_uses_authoritative = False + raise self._async_bwd_qlen = qlen self._async_bwd_output_device = output_device + self._async_bwd_uses_authoritative = use_authoritative def sync_backward(self) -> Tuple[torch.Tensor, torch.Tensor]: """Wait for async backward and return results.""" - self.cpu_infer.sync() - - qlen = self._async_bwd_qlen - output_device = self._async_bwd_output_device - buffer = self._get_buffer(qlen) + if not hasattr(self, "_async_bwd_qlen") or self._async_bwd_qlen is None: + raise RuntimeError("No pending backward. Call submit_backward_async() first.") + + use_authoritative = getattr(self, "_async_bwd_uses_authoritative", False) + try: + self.cpu_infer.sync() + qlen = self._async_bwd_qlen + output_device = self._async_bwd_output_device + buffer = self._get_buffer(qlen) + result = self._return_grads(buffer, qlen, output_device) + if use_authoritative: + self._publish_authoritative_optimizer_grads() + except Exception: + if use_authoritative: + self._abort_authoritative_optimizer_grad_write() + self._cache_depth = max(0, self._cache_depth - 1) + self._async_bwd_qlen = None + self._async_bwd_output_device = None + self._async_bwd_uses_authoritative = False + raise self._cache_depth -= 1 - return self._return_grads(buffer, qlen, output_device) + self._async_bwd_qlen = None + self._async_bwd_output_device = None + self._async_bwd_uses_authoritative = False + return result # ========== Backward repack (optional, subclasses may override) ========== diff --git a/kt-kernel/python/sft/dist_utils.py b/kt-kernel/python/sft/dist_utils.py index 391ba536b..50f021405 100644 --- a/kt-kernel/python/sft/dist_utils.py +++ b/kt-kernel/python/sft/dist_utils.py @@ -10,11 +10,43 @@ from __future__ import annotations from contextlib import nullcontext +import os from typing import Any import torch +def _distributed_rank_world_size() -> tuple[int, int]: + """Return rank/world size during both model construction and runtime. + + Accelerate/torchrun may construct the model before the process group is + initialized. In that phase the standard launcher environment is the only + reliable way to keep rank-0 KT ownership and rank-0 buffer capacity + consistent across processes. + """ + import torch.distributed as dist + + if dist.is_initialized(): + return int(dist.get_rank()), int(dist.get_world_size()) + + rank_text = os.environ.get("RANK") + world_text = os.environ.get("WORLD_SIZE") + if rank_text is None or world_text is None: + return 0, 1 + try: + rank = int(rank_text) + world_size = int(world_text) + except ValueError as exc: + raise RuntimeError( + f"Invalid distributed launcher environment: RANK={rank_text!r}, WORLD_SIZE={world_text!r}" + ) from exc + if world_size <= 0 or rank < 0 or rank >= world_size: + raise RuntimeError( + f"Invalid distributed launcher environment: rank={rank}, world_size={world_size}" + ) + return rank, world_size + + def _all_gather_qlens(local_qlen: int, device: torch.device, world_size: int) -> list[int]: import torch.distributed as dist diff --git a/kt-kernel/python/sft/layer.py b/kt-kernel/python/sft/layer.py index 4f0acd902..9b7bd339c 100644 --- a/kt-kernel/python/sft/layer.py +++ b/kt-kernel/python/sft/layer.py @@ -46,6 +46,8 @@ def __init__( hidden_size: int, layer_idx: int, lora_experts: "LoRAExperts | None" = None, + full_weight_grad: bool | None = None, + uses_authoritative_optimizer_grads: bool | None = None, ): super().__init__() self._is_kt_moe_wrapper = True @@ -106,9 +108,18 @@ def __init__( # _peft_lora_modules: {expert_idx: {proj_name: (lora_A, lora_B)}} self._peft_lora_modules: dict[int, dict[str, tuple[nn.Module, nn.Module]]] | None = None self._lora_pointers_dirty = False - - # Full weight grad mode (set during wrapping or kt_adapt_peft_lora) - self._full_weight_grad = getattr(wrapper, "_full_weight_grad", False) if wrapper is not None else False + self._kt_managed_lora_enabled = False + + # Training-mode flags must be identical on every distributed rank even + # though only rank 0 owns the backend object. + if full_weight_grad is None: + full_weight_grad = getattr(wrapper, "_full_weight_grad", False) if wrapper is not None else False + if uses_authoritative_optimizer_grads is None: + uses_authoritative_optimizer_grads = bool( + wrapper is not None and getattr(wrapper, "_uses_authoritative_optimizer_grads", False) + ) + self._full_weight_grad = bool(full_weight_grad) + self._uses_authoritative_optimizer_grads = bool(uses_authoritative_optimizer_grads) def _apply(self, fn, recurse=True): # Protect experts from device transfer (PEFT LoRA should stay on CPU for KT) @@ -139,7 +150,11 @@ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: with torch.profiler.record_function("kt.sft.routing"): topk_ids, topk_weights = self._compute_routing(hidden_states) - train_lora = self._peft_lora_modules is not None and len(self._peft_lora_modules) > 0 + train_lora = bool( + self._kt_managed_lora_enabled + or (self._peft_lora_modules is not None and len(self._peft_lora_modules) > 0) + or getattr(self, "_fused_expert_lora_params", None) + ) full_weight_grad = self._full_weight_grad save_for_backward = ( @@ -186,15 +201,21 @@ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: # Use KTMoEFunction whenever backward is needed so KT backward and LoRA # gradient paths remain connected. if use_autograd_path: - lora_ref = hidden_states.new_empty(()) + # A requires-grad sentinel keeps the custom autograd node alive on + # non-rank-0 fused/full ranks that intentionally own no KT params. + lora_ref = hidden_states.new_empty((), requires_grad=(train_lora or full_weight_grad)) if train_lora and self._peft_lora_modules: + found_lora_ref = False for expert_loras in self._peft_lora_modules.values(): for lora_A, lora_B in expert_loras.values(): if hasattr(lora_A, "weight") and lora_A.weight.requires_grad: lora_ref = lora_A.weight + found_lora_ref = True break - if lora_ref.numel() > 0: + if found_lora_ref: break + elif train_lora and getattr(self, "_fused_expert_lora_params", None): + lora_ref = self._fused_expert_lora_params[0] elif full_weight_grad and self.wrapper is not None: # In full mode, use base weight param as autograd sentinel if self.wrapper.gate_proj_buf is not None: @@ -404,6 +425,20 @@ def _submit_and_compute_gpu( rank = dist.get_rank() if dist.is_initialized() else 0 world_size = dist.get_world_size() if dist_on else 1 + if dist_on and self._uses_authoritative_optimizer_grads: + wrapped_world_size = int(getattr(self, "_kt_world_size_at_wrap", world_size)) + if wrapped_world_size != world_size: + raise RuntimeError( + f"Layer {self.layer_idx}: KT wrapper was created for world_size={wrapped_world_size}, " + f"but the active process group has world_size={world_size}" + ) + if rank == 0 and self.wrapper is None: + raise RuntimeError(f"Layer {self.layer_idx}: rank 0 does not own the authoritative KT backend") + if rank != 0 and self.wrapper is not None: + raise RuntimeError( + f"Layer {self.layer_idx}: rank {rank} unexpectedly owns an authoritative KT backend" + ) + qlen = batch_size * seq_len if dist_on: diff --git a/kt-kernel/python/sft/lora.py b/kt-kernel/python/sft/lora.py index e9f8d6202..b90fc5d15 100644 --- a/kt-kernel/python/sft/lora.py +++ b/kt-kernel/python/sft/lora.py @@ -22,6 +22,7 @@ import torch.nn as nn from .arch import MOEArchConfig +from .dist_utils import _distributed_rank_world_size logger = logging.getLogger(__name__) @@ -98,25 +99,30 @@ def _find_kt_wrappers(model: nn.Module): return wrappers +def _collect_wrapper_managed_lora_params(wrapper) -> list[nn.Parameter]: + """Collect only C++-managed PEFT/fused LoRA params for one wrapper.""" + params: list[nn.Parameter] = [] + peft_lora_modules = getattr(wrapper, "_peft_lora_modules", None) + if peft_lora_modules is not None: + for expert_loras in peft_lora_modules.values(): + for lora_A, lora_B in expert_loras.values(): + if hasattr(lora_A, "weight") and lora_A.weight.requires_grad: + params.append(lora_A.weight) + if hasattr(lora_B, "weight") and lora_B.weight.requires_grad: + params.append(lora_B.weight) + fused_params = getattr(wrapper, "_fused_expert_lora_params", None) + if fused_params is not None: + params.extend(fused_params) + return params + + def _collect_kt_lora_params(wrappers) -> list[nn.Parameter]: """Collect LoRA-only trainable parameters from KT wrappers.""" params: list[nn.Parameter] = [] if wrappers: for wrapper in wrappers: - # PEFT LoRA parameters (from _peft_lora_modules) - peft_lora_modules = getattr(wrapper, "_peft_lora_modules", None) - if peft_lora_modules is not None: - for expert_loras in peft_lora_modules.values(): - for lora_A, lora_B in expert_loras.values(): - if hasattr(lora_A, "weight") and lora_A.weight.requires_grad: - params.append(lora_A.weight) - if hasattr(lora_B, "weight") and lora_B.weight.requires_grad: - params.append(lora_B.weight) - # Fused expert LoRA parameters (KT-managed, not PEFT) - fused_params = getattr(wrapper, "_fused_expert_lora_params", None) - if fused_params is not None: - params.extend(fused_params) + params.extend(_collect_wrapper_managed_lora_params(wrapper)) # lora_experts parameters (separate feature) if getattr(wrapper, "lora_experts", None) is not None: params.extend(wrapper.lora_experts.parameters()) @@ -198,17 +204,14 @@ def kt_adapt_peft_lora(model: nn.Module) -> None: Should be called after PEFT LoRA injection and before create_optimizer. """ - import torch.distributed as dist - wrappers = _find_kt_wrappers(model) if not wrappers: logger.info("[kt_adapt_peft_lora] No _kt_wrappers found, skipping") return - is_rank_0 = True - if dist.is_initialized(): - is_rank_0 = dist.get_rank() == 0 + distributed_rank, _ = _distributed_rank_world_size() + is_rank_0 = distributed_rank == 0 adapted_count = 0 for wrapper in wrappers: @@ -221,10 +224,11 @@ def kt_adapt_peft_lora(model: nn.Module) -> None: continue # Fused experts (transformers v5): PEFT cannot auto-attach LoRA to packed - # nn.Parameter tensors. Create KT-managed LoRA buffers with proper init, - # wrap as nn.Parameter for optimizer, and pre-assign .grad for C++ backward. + # nn.Parameter tensors. Create KT-managed LoRA buffers with proper init + # and wrap them as nn.Parameter objects for optimizer injection. if getattr(wrapper, "_fused_experts", False): lora_rank = getattr(wrapper, "_lora_rank", 1) + authoritative_mode = bool(getattr(wrapper, "_uses_authoritative_optimizer_grads", False)) # In full mode (lora_rank=0), skip LoRA buffer creation entirely. # C++ kernel will not compute LoRA contributions when lora_rank=0. @@ -238,11 +242,23 @@ def kt_adapt_peft_lora(model: nn.Module) -> None: adapted_count += 1 continue + wrapper._kt_managed_lora_enabled = True + + # In rank-0-authoritative BF16 mode, non-rank-0 processes need the + # mode flag for collective/autograd symmetry but own no optimizer + # Parameter or gradient buffer. + if authoritative_mode and not is_rank_0: + wrapper._fused_expert_lora_params = [] + wrapper._peft_lora_modules = None + adapted_count += 1 + continue + lora_buffers, lora_grad_buffers, lora_params = _create_fused_expert_lora_buffers( wrapper, moe_config, lora_rank, torch.bfloat16, + preassign_grads=not authoritative_mode, ) if is_rank_0 and wrapper.wrapper is not None: @@ -254,6 +270,16 @@ def kt_adapt_peft_lora(model: nn.Module) -> None: f"[kt_adapt_peft_lora] Layer {layer_idx}: fused expert LoRA " f"(r={lora_rank}, E={moe_config.expert_num})" ) + if authoritative_mode: + for key, param in zip( + ("gate_lora_a", "gate_lora_b", "up_lora_a", "up_lora_b", "down_lora_a", "down_lora_b"), + lora_params, + ): + wrapper.wrapper.register_authoritative_optimizer_grad( + f"lora.{key}", param, lora_grad_buffers[f"grad_{key}"] + ) + elif authoritative_mode: + raise RuntimeError(f"Layer {layer_idx}: rank 0 authoritative LoRA requires a KT backend") wrapper._fused_expert_lora_params = lora_params wrapper._peft_lora_modules = None @@ -306,13 +332,14 @@ def kt_adapt_peft_lora(model: nn.Module) -> None: # Store PEFT LoRA references on wrapper wrapper._peft_lora_modules = peft_lora_modules - # In full_weight_grad mode, PEFT LoRA is not injected by LlamaFactory, - # so no PEFT LoRA found is expected — skip the error. + # Missing PEFT LoRA is valid only for pure Full. Hybrid has lora_rank + # greater than zero and must fail instead of silently training base + # weights alone. if not peft_lora_modules: - if getattr(wrapper, "_full_weight_grad", False): + if getattr(wrapper, "_full_weight_grad", False) and getattr(wrapper, "_lora_rank", 0) == 0: logger.info( f"[kt_adapt_peft_lora] Layer {layer_idx}: No PEFT LoRA found " - f"(full_weight_grad mode — expected, skipping)" + f"(pure Full mode — expected, skipping)" ) adapted_count += 1 continue @@ -321,6 +348,9 @@ def kt_adapt_peft_lora(model: nn.Module) -> None: f"Check that PEFT lora_target includes expert modules." ) + wrapper._kt_managed_lora_enabled = True + authoritative_mode = bool(getattr(wrapper, "_uses_authoritative_optimizer_grads", False)) + # Allocate contiguous bf16 buffers and populate with initial PEFT values (all ranks) lora_buffers = _create_lora_view_buffers(peft_lora_modules, moe_config, torch.bfloat16) lora_grad_buffers = _create_lora_grad_buffers(peft_lora_modules, moe_config) @@ -333,7 +363,14 @@ def kt_adapt_peft_lora(model: nn.Module) -> None: logger.info(f"[kt_adapt_peft_lora] Layer {layer_idx}: synced PEFT LoRA to C++ kernel") # All ranks: replace PEFT weights with views into the contiguous buffers - _replace_peft_weights_with_views(peft_lora_modules, lora_buffers, lora_grad_buffers, moe_config) + _replace_peft_weights_with_views( + peft_lora_modules, + lora_buffers, + lora_grad_buffers, + moe_config, + authoritative_mode=authoritative_mode, + authoritative_backend=wrapper.wrapper if authoritative_mode and is_rank_0 else None, + ) adapted_count += 1 @@ -487,6 +524,7 @@ def _create_fused_expert_lora_buffers( moe_config: MOEArchConfig, lora_rank: int, dtype: torch.dtype = torch.bfloat16, + preassign_grads: bool = True, ) -> tuple[dict[str, torch.Tensor], dict[str, torch.Tensor], list[nn.Parameter]]: """ Create KT-managed LoRA buffers for fused expert modules. @@ -532,7 +570,8 @@ def _create_fused_expert_lora_buffers( lora_params = [] for key in ("gate_lora_a", "gate_lora_b", "up_lora_a", "up_lora_b", "down_lora_a", "down_lora_b"): param = nn.Parameter(lora_buffers[key], requires_grad=True) - param.grad = lora_grad_buffers[f"grad_{key}"] + if preassign_grads: + param.grad = lora_grad_buffers[f"grad_{key}"] lora_params.append(param) return lora_buffers, lora_grad_buffers, lora_params @@ -548,6 +587,9 @@ def _replace_peft_weights_with_views( buffers: dict[str, torch.Tensor], grad_buffers: dict[str, torch.Tensor], moe_config: MOEArchConfig, + *, + authoritative_mode: bool = False, + authoritative_backend=None, ) -> None: """ Replace each PEFT LoRA module's .weight with a view into the contiguous buffer. @@ -585,8 +627,21 @@ def _replace_peft_weights_with_views( lora_B.weight.data = buffers[key_b][expert_idx] lora_A.weight.requires_grad_(True) lora_B.weight.requires_grad_(True) - lora_A.weight.grad = grad_buffers["grad_" + key_a][expert_idx] - lora_B.weight.grad = grad_buffers["grad_" + key_b][expert_idx] + grad_view_a = grad_buffers["grad_" + key_a][expert_idx] + grad_view_b = grad_buffers["grad_" + key_b][expert_idx] + if authoritative_mode: + lora_A.weight.grad = None + lora_B.weight.grad = None + if authoritative_backend is not None: + authoritative_backend.register_authoritative_optimizer_grad( + f"lora.{key_a}.expert_{expert_idx}", lora_A.weight, grad_view_a + ) + authoritative_backend.register_authoritative_optimizer_grad( + f"lora.{key_b}.expert_{expert_idx}", lora_B.weight, grad_view_b + ) + else: + lora_A.weight.grad = grad_view_a + lora_B.weight.grad = grad_view_b if not _first_logged: _new_id_a = id(lora_A.weight) @@ -625,10 +680,14 @@ def update_kt_lora_pointers(model: nn.Module): if wrappers: for wrapper in wrappers: - wrapper._lora_pointers_dirty = True + if getattr(wrapper, "_kt_managed_lora_enabled", False): + wrapper._lora_pointers_dirty = True # In full mode, base weights also need re-sync after optimizer step if getattr(wrapper, "_full_weight_grad", False) and wrapper.wrapper is not None: wrapper.wrapper._base_weights_dirty = True + backend = getattr(wrapper, "wrapper", None) + if backend is not None and getattr(backend, "_uses_authoritative_optimizer_grads", False): + backend.release_authoritative_optimizer_grads() # ============================================================================= @@ -649,11 +708,48 @@ def sync_kt_lora_gradients(model: nn.Module) -> None: return world_size = dist.get_world_size() + rank = dist.get_rank() if world_size <= 1: return - # Sync base weight gradients in full mode wrappers = _find_kt_wrappers(model) + if not wrappers: + return + + # AMXBF16_SFT gathers every rank's routed rows and writes a world-size + # normalized optimizer gradient on rank 0. No gradient collective is + # needed here; validate ownership/aliases only. Ordinary GPU + # lora_experts are deliberately excluded and remain DDP/FSDP-managed. + authoritative_wrappers = [w for w in wrappers if getattr(w, "_uses_authoritative_optimizer_grads", False)] + if authoritative_wrappers: + if len(authoritative_wrappers) != len(wrappers): + raise RuntimeError("Mixed authoritative and legacy KT SFT backends are unsupported in one model") + for wrapper in authoritative_wrappers: + backend = getattr(wrapper, "wrapper", None) + wrapped_world_size = int(getattr(wrapper, "_kt_world_size_at_wrap", world_size)) + if wrapped_world_size != world_size: + raise RuntimeError( + f"Layer {wrapper.layer_idx}: KT wrapper was created for world_size={wrapped_world_size}, " + f"but the active process group has world_size={world_size}" + ) + if rank == 0: + if backend is None: + raise RuntimeError(f"Layer {wrapper.layer_idx}: rank 0 does not own the authoritative KT backend") + backend.validate_authoritative_optimizer_grad_state() + else: + if backend is not None: + raise RuntimeError( + f"Layer {wrapper.layer_idx}: rank {rank} unexpectedly owns an authoritative KT backend" + ) + for param in _collect_wrapper_managed_lora_params(wrapper): + if param.grad is not None: + raise RuntimeError( + f"Layer {wrapper.layer_idx}: non-rank-0 KT LoRA Parameter unexpectedly has a gradient" + ) + return + + # Legacy backends retain their existing cross-rank synchronization. + # Sync base weight gradients in full mode. if wrappers: for wrapper in wrappers: if not getattr(wrapper, "_full_weight_grad", False): diff --git a/kt-kernel/python/sft/wrapper.py b/kt-kernel/python/sft/wrapper.py index 4e3dd4614..c727802e1 100644 --- a/kt-kernel/python/sft/wrapper.py +++ b/kt-kernel/python/sft/wrapper.py @@ -23,6 +23,8 @@ ) from .layer import KTMoELayerWrapper from .lora import LoRAExperts +from .base import _supports_authoritative_optimizer_grads +from .dist_utils import _distributed_rank_world_size from .weights import ( _clear_original_expert_weights, extract_moe_weights, @@ -144,15 +146,13 @@ def wrap_moe_layers_with_kt_wrapper(model: nn.Module, kt_plugin: Any) -> list[KT LoRA is handled by PEFT and later adapted via kt_adapt_peft_lora(). Only rank 0 initializes KT kernel and loads weights. """ - import torch.distributed as dist - if not KT_KERNEL_AVAILABLE: raise KTAMXNotAvailableError("kt_kernel not found. Please install kt_kernel to enable KT MoE support.") - # Only rank 0 should initialize KT and load weights - is_rank_0 = True - if dist.is_initialized(): - is_rank_0 = dist.get_rank() == 0 + # Only global rank 0 initializes KT. Launcher env fallback matters when + # model construction happens before init_process_group(). + distributed_rank, distributed_world_size = _distributed_rank_world_size() + is_rank_0 = distributed_rank == 0 moe_config = get_moe_arch_config(model.config) _text_cfg = getattr(model.config, "text_config", model.config) @@ -171,11 +171,14 @@ def wrap_moe_layers_with_kt_wrapper(model: nn.Module, kt_plugin: Any) -> list[KT # Read full_weight_grad mode _raw_fwg = getattr(cfg, "kt_full_weight_grad", None) full_weight_grad = _raw_fwg if _raw_fwg is not None else False - - # In full mode, lora_rank should be 0 (no LoRA, only base weight grad) - # If user explicitly set lora_rank > 0 in full mode (hybrid), keep it. - # Otherwise, auto-set lora_rank=0. - if full_weight_grad and lora_rank > 0: + train_mode = getattr(cfg, "kt_train_mode", "lora") + + # Full and hybrid are explicit modes. LlamaFactory exposes a default + # lora_rank even for full tuning, which must not silently turn Full into + # Hybrid. Preserve the legacy fallback for callers without train_mode. + if train_mode == "full": + lora_rank = 0 + elif full_weight_grad and train_mode != "hybrid" and lora_rank > 0: _has_explicit_lora_rank = getattr(cfg, "kt_lora_rank", None) is not None if not _has_explicit_lora_rank: lora_rank = 0 @@ -217,6 +220,10 @@ def wrap_moe_layers_with_kt_wrapper(model: nn.Module, kt_plugin: Any) -> list[KT if "SkipLoRA" in kt_method: logger.info(f"Using SkipLoRA backend: {kt_method} (MoE LoRA gradients will be skipped)") + requested_num_gpu_experts = int(getattr(cfg, "kt_num_gpu_experts", 0) or 0) + uses_authoritative_optimizer_grads = _supports_authoritative_optimizer_grads( + kt_method, requested_num_gpu_experts + ) threadpool_count = getattr(cfg, "kt_threadpool_count", 1) if getattr(cfg, "kt_tp_enabled", False) else 1 @@ -277,10 +284,6 @@ def wrap_moe_layers_with_kt_wrapper(model: nn.Module, kt_plugin: Any) -> list[KT "files could be resolved for on-the-fly expert loading." ) - import torch.distributed as _dist - - _rank = _dist.get_rank() if _dist.is_initialized() else 0 - model_container, layers = _get_model_container_and_layers(model, purpose="wrapping") logger.info(f"Total layers={len(layers)}, is_rank_0={is_rank_0}") @@ -328,6 +331,10 @@ def wrap_moe_layers_with_kt_wrapper(model: nn.Module, kt_plugin: Any) -> list[KT chunked_prefill_size = getattr(cfg, "kt_model_max_length", None) if chunked_prefill_size is None: chunked_prefill_size = getattr(model.config, "max_position_embeddings", 4096) + # Rank 0 receives the concatenation of every rank's local rows. Model + # configs are homogeneous across ranks, so the sum of local maxima is + # the per-rank capacity multiplied by world size. + rank0_chunked_prefill_size = int(chunked_prefill_size) * distributed_world_size # Only rank 0 creates KTMoEWrapper and loads weights if is_rank_0: @@ -342,7 +349,7 @@ def wrap_moe_layers_with_kt_wrapper(model: nn.Module, kt_plugin: Any) -> list[KT cpuinfer_threads=getattr(cfg, "kt_num_threads", 1), threadpool_count=threadpool_count, weight_path=kt_weight_path or "", - chunked_prefill_size=chunked_prefill_size, + chunked_prefill_size=rank0_chunked_prefill_size, method=kt_method, mode="sft", lora_rank=lora_rank, @@ -350,10 +357,15 @@ def wrap_moe_layers_with_kt_wrapper(model: nn.Module, kt_plugin: Any) -> list[KT max_cache_depth=getattr(cfg, "kt_max_cache_depth", 2), full_weight_grad=full_weight_grad, ) + # The current SFT wrapping path routes all experts through KT even + # when the loading config requested GPU experts. Preserve that + # configuration's legacy gradient lifecycle until the hybrid + # routed-expert path supports authoritative buffers end to end. + wrapper._uses_authoritative_optimizer_grads = uses_authoritative_optimizer_grads # Set share_backward_bb and share_cache_pool BEFORE load_weights (config is built during load) wrapper.share_backward_bb = cfg.kt_share_backward_bb - single_process = not dist.is_initialized() or dist.get_world_size() == 1 + single_process = distributed_world_size == 1 reuse_checkpoint_forward = ( single_process and full_weight_grad @@ -415,9 +427,13 @@ def wrap_moe_layers_with_kt_wrapper(model: nn.Module, kt_plugin: Any) -> list[KT hidden_size=hidden_size, layer_idx=layer_idx, lora_experts=lora_experts, + full_weight_grad=full_weight_grad, + uses_authoritative_optimizer_grads=uses_authoritative_optimizer_grads, ) layer_wrapper._fused_experts = _layer_is_fused layer_wrapper._lora_rank = lora_rank + layer_wrapper._kt_owner_rank = 0 + layer_wrapper._kt_world_size_at_wrap = distributed_world_size setattr(layer, moe_config.moe_layer_attr, layer_wrapper) # Base weights have been copied into the C++ kernel's internal BufferB format. @@ -473,6 +489,12 @@ class should import it from the appropriate dataclasses module. } kt_train_mode = kt_train_mode_map.get(finetuning_type, None) if finetuning_type else None + configured_lora_rank = getattr(finetuning_args, "lora_rank", None) if finetuning_args else None + configured_lora_alpha = getattr(finetuning_args, "lora_alpha", None) if finetuning_args else None + if kt_train_mode == "full": + configured_lora_rank = None + configured_lora_alpha = None + kt_config = KTConfig( kt_backend=getattr(model_args, "kt_backend", None), kt_num_threads=getattr(model_args, "kt_num_threads", None), @@ -485,8 +507,8 @@ class should import it from the appropriate dataclasses module. kt_use_lora_experts=getattr(model_args, "kt_use_lora_experts", None), kt_lora_expert_num=getattr(model_args, "kt_lora_expert_num", None), kt_lora_expert_intermediate_size=getattr(model_args, "kt_lora_expert_intermediate_size", None), - kt_lora_rank=getattr(finetuning_args, "lora_rank", None) if finetuning_args else None, - kt_lora_alpha=getattr(finetuning_args, "lora_alpha", None) if finetuning_args else None, + kt_lora_rank=configured_lora_rank, + kt_lora_alpha=configured_lora_alpha, kt_model_max_length=getattr(model_args, "model_max_length", None), kt_train_mode=kt_train_mode, ) @@ -648,6 +670,8 @@ def load_kt_model( model._kt_wrappers = wrappers model._kt_tp_enabled = bool(getattr(cfg, "kt_tp_enabled", False)) model._kt_use_lora_experts = bool(getattr(cfg, "kt_use_lora_experts", False)) + model._kt_full_weight_grad = bool(getattr(cfg, "kt_full_weight_grad", False)) + model._kt_train_mode = getattr(cfg, "kt_train_mode", "lora") logger.info("Model loaded with KTMoEWrapper backend successfully") return model diff --git a/kt-kernel/test/per_commit/test_sft_authoritative_grad.py b/kt-kernel/test/per_commit/test_sft_authoritative_grad.py new file mode 100644 index 000000000..a48d67818 --- /dev/null +++ b/kt-kernel/test/per_commit/test_sft_authoritative_grad.py @@ -0,0 +1,361 @@ +# SPDX-License-Identifier: Apache-2.0 + +import os +from types import SimpleNamespace + +import pytest +import torch + +from kt_kernel.sft.base import BaseSFTMoEWrapper, _supports_authoritative_optimizer_grads +from kt_kernel.sft.autograd import KTMoEFunction +from kt_kernel.sft.dist_utils import _distributed_rank_world_size +from kt_kernel.sft.lora import kt_adapt_peft_lora, update_kt_lora_pointers + + +class _TaskRunner: + def __init__(self): + self.pending = None + self.fail_next_submit = False + self.fail_next_sync = False + + def submit(self, task): + if self.fail_next_submit: + self.fail_next_submit = False + raise RuntimeError("synthetic submit failure") + if self.pending is not None: + raise RuntimeError("task already pending") + self.pending = task + + def sync(self): + task = self.pending + self.pending = None + if self.fail_next_sync: + self.fail_next_sync = False + raise RuntimeError("synthetic C++ failure") + if task is not None: + task() + + +def test_capability_is_limited_to_cpu_only_amxbf16_sft(): + assert _supports_authoritative_optimizer_grads("AMXBF16_SFT", 0) + assert not _supports_authoritative_optimizer_grads("AMXBF16_SFT", 1) + assert not _supports_authoritative_optimizer_grads("AMXINT8_SFT", 0) + assert not _supports_authoritative_optimizer_grads("AMXINT4_SFT", 0) + assert not _supports_authoritative_optimizer_grads("AMXBF16_SFT_SkipLoRA", 0) + + +class _FakeAuthoritativeWrapper(BaseSFTMoEWrapper): + """Minimal backend exercising BaseSFTMoEWrapper's real lifecycle.""" + + def __init__(self, parameter_count=1): + # Avoid constructing CPUInfer or importing a real AMX extension. + self._uses_authoritative_optimizer_grads = True + self._init_authoritative_optimizer_grads() + self._cache_depth = 1 + self._base_weights_dirty = False + self.cpu_infer = _TaskRunner() + self.buffer = SimpleNamespace() + self.write_value = 2.0 + self.task_modes = [] + self.fail_return_grads = False + self.staging_copy_count = 0 + self.parameters = [] + self.grad_views = [] + self._full_weight_grad = True + self.share_backward_bb = False + for idx in range(parameter_count): + parameter = torch.nn.Parameter(torch.ones(4, dtype=torch.float32)) + grad_view = torch.full_like(parameter, -99.0) + self.parameters.append(parameter) + self.grad_views.append(grad_view) + self.register_authoritative_optimizer_grad(f"fake.{idx}", parameter, grad_view) + + def _get_buffer(self, _qlen): + return self.buffer + + def _copy_grad_output_to_cpu(self, _buffer, _grad_output, _qlen): + self.staging_copy_count += 1 + return None + + def _return_grads(self, _buffer, qlen, _output_device): + if self.fail_return_grads: + raise RuntimeError("synthetic return failure") + return torch.zeros(qlen, 1), torch.zeros(qlen, 1) + + def sync_forward(self, output_device=None): + output = torch.zeros(1, 1) + return output if output_device is None else output.to(output_device) + + def clear_checkpoint_output(self): + return None + + def _make_forward_task(self, _buffer, _save_for_backward): + raise NotImplementedError + + def _make_backward_task( + self, + _buffer, + accumulate_optimizer_grads=False, + optimizer_grad_scale=1.0, + ): + self.task_modes.append((bool(accumulate_optimizer_grads), float(optimizer_grad_scale))) + + def task(): + value = self.write_value * float(optimizer_grad_scale) + for grad_view in self.grad_views: + if accumulate_optimizer_grads: + grad_view.add_(value) + else: + grad_view.fill_(value) + + return task + + def reset_cache(self): + self._cache_depth = 1 + + # Abstract backend hooks not needed by these tests. + def load_weights(self, physical_to_logical_map_cpu): + raise NotImplementedError + + def init_lora_weights(self, *args, **kwargs): + raise NotImplementedError + + def update_lora_weights(self): + raise NotImplementedError + + def update_base_weights(self): + raise NotImplementedError + + +def test_sync_backward_overwrite_accumulate_publish_and_step_release(): + backend = _FakeAuthoritativeWrapper() + parameter = backend.parameters[0] + grad_view = backend.grad_views[0] + + assert parameter.grad is None + backend.backward(torch.ones(1, 1), optimizer_grad_scale=0.5) + assert backend.task_modes == [(False, 0.5)] + assert parameter.grad is grad_view + torch.testing.assert_close(grad_view, torch.ones_like(grad_view)) + + backend.reset_cache() + backend.write_value = 4.0 + backend.backward(torch.ones(1, 1), optimizer_grad_scale=0.5) + assert backend.task_modes[-1] == (True, 0.5) + assert parameter.grad is grad_view + torch.testing.assert_close(grad_view, torch.full_like(grad_view, 3.0)) + + optimizer = torch.optim.SGD([parameter], lr=0.1) + parameter_before_step = parameter.detach().clone() + optimizer.step() + torch.testing.assert_close(parameter, parameter_before_step - 0.1 * grad_view) + + layer = SimpleNamespace( + layer_idx=0, + wrapper=backend, + _kt_managed_lora_enabled=True, + _lora_pointers_dirty=False, + _full_weight_grad=True, + ) + update_kt_lora_pointers(SimpleNamespace(_kt_wrappers=[layer])) + assert layer._lora_pointers_dirty + assert backend._base_weights_dirty + assert parameter.grad is None + + raw_before_zero_grad = grad_view.clone() + optimizer.zero_grad(set_to_none=False) + assert parameter.grad is None + torch.testing.assert_close(grad_view, raw_before_zero_grad) + + backend.reset_cache() + backend.backward(torch.ones(1, 1)) + assert backend.task_modes[-1] == (False, 1.0) + + +def test_mixed_foreign_and_changed_metadata_fail_fast(): + backend = _FakeAuthoritativeWrapper(parameter_count=2) + backend.parameters[0].grad = backend.grad_views[0] + with pytest.raises(RuntimeError, match="Mixed authoritative"): + backend._prepare_authoritative_optimizer_grad_write(1.0) + + backend.parameters[1].grad = backend.grad_views[1] + backend.parameters[0].grad = backend.grad_views[0].view_as(backend.grad_views[0]) + with pytest.raises(RuntimeError, match="externally replaced"): + backend._prepare_authoritative_optimizer_grad_write(1.0) + + metadata_backend = _FakeAuthoritativeWrapper() + metadata_backend.grad_views[0].data = torch.zeros(4, dtype=torch.float32) + with pytest.raises(RuntimeError, match="metadata changed"): + metadata_backend._prepare_authoritative_optimizer_grad_write(1.0) + + +def test_failed_sync_closes_window_and_retry_overwrites(): + backend = _FakeAuthoritativeWrapper() + backend.cpu_infer.fail_next_sync = True + + with pytest.raises(RuntimeError, match=r"synthetic C\+\+ failure"): + backend.backward(torch.ones(1, 1)) + assert backend.parameters[0].grad is None + assert backend.validate_authoritative_optimizer_grad_state() == "closed" + assert backend._cache_depth == 0 + + backend.reset_cache() + backend.backward(torch.ones(1, 1)) + assert backend.task_modes[-1] == (False, 1.0) + assert backend.parameters[0].grad is backend.grad_views[0] + + +def test_post_cpp_return_failure_closes_sync_and_async_windows(): + backend = _FakeAuthoritativeWrapper() + backend.fail_return_grads = True + + with pytest.raises(RuntimeError, match="synthetic return failure"): + backend.backward(torch.ones(1, 1)) + assert backend.parameters[0].grad is None + assert backend._cache_depth == 0 + + backend.reset_cache() + backend.submit_backward_async(torch.ones(1, 1)) + with pytest.raises(RuntimeError, match="synthetic return failure"): + backend.sync_backward() + assert backend.parameters[0].grad is None + assert backend._cache_depth == 0 + assert backend._async_bwd_qlen is None + + backend.fail_return_grads = False + backend.reset_cache() + backend.backward(torch.ones(1, 1)) + assert backend.task_modes[-1] == (False, 1.0) + + +def test_async_submit_failure_invalidates_cache_and_pending_state(): + backend = _FakeAuthoritativeWrapper() + backend.cpu_infer.fail_next_submit = True + + with pytest.raises(RuntimeError, match="synthetic submit failure"): + backend.submit_backward_async(torch.ones(1, 1)) + assert backend.parameters[0].grad is None + assert backend._cache_depth == 0 + assert backend._async_bwd_qlen is None + + +def test_pending_async_backward_rejects_reentrant_submit_before_staging_copy(): + backend = _FakeAuthoritativeWrapper() + backend.submit_backward_async(torch.ones(1, 1)) + + with pytest.raises(RuntimeError, match="already pending"): + backend.submit_backward_async(torch.full((1, 1), 7.0)) + with pytest.raises(RuntimeError, match="already pending"): + backend.backward(torch.full((1, 1), 9.0)) + assert backend.staging_copy_count == 1 + + backend.sync_backward() + assert backend.parameters[0].grad is backend.grad_views[0] + torch.testing.assert_close(backend.grad_views[0], torch.full_like(backend.grad_views[0], 2.0)) + + +def test_hybrid_requires_expert_lora_instead_of_silently_falling_back_to_full(): + class _Expert(torch.nn.Module): + def __init__(self): + super().__init__() + self.gate_proj = torch.nn.Linear(4, 4, bias=False) + self.up_proj = torch.nn.Linear(4, 4, bias=False) + self.down_proj = torch.nn.Linear(4, 4, bias=False) + + layer = SimpleNamespace( + layer_idx=0, + moe_config=SimpleNamespace(weight_names=("gate_proj", "up_proj", "down_proj")), + experts=torch.nn.ModuleList([_Expert()]), + _experts_attr="experts", + _fused_experts=False, + _lora_rank=4, + _full_weight_grad=True, + ) + with pytest.raises(RuntimeError, match="No PEFT LoRA found"): + kt_adapt_peft_lora(SimpleNamespace(_kt_wrappers=[layer])) + + +def test_launcher_environment_preserves_rank0_ownership_before_process_group_init(): + previous_rank = os.environ.get("RANK") + previous_world = os.environ.get("WORLD_SIZE") + try: + os.environ["RANK"] = "1" + os.environ["WORLD_SIZE"] = "2" + assert _distributed_rank_world_size() == (1, 2) + finally: + if previous_rank is None: + os.environ.pop("RANK", None) + else: + os.environ["RANK"] = previous_rank + if previous_world is None: + os.environ.pop("WORLD_SIZE", None) + else: + os.environ["WORLD_SIZE"] = previous_world + + +def test_async_backward_publishes_only_after_successful_sync(): + backend = _FakeAuthoritativeWrapper() + parameter = backend.parameters[0] + + backend.submit_backward_async(torch.ones(1, 1), optimizer_grad_scale=0.25) + assert backend.task_modes == [(False, 0.25)] + assert parameter.grad is None + backend.sync_backward() + assert parameter.grad is backend.grad_views[0] + + backend.reset_cache() + backend.write_value = 8.0 + backend.submit_backward_async(torch.ones(1, 1), optimizer_grad_scale=0.25) + assert backend.task_modes[-1] == (True, 0.25) + assert parameter.grad is backend.grad_views[0] + backend.sync_backward() + torch.testing.assert_close(backend.grad_views[0], torch.full_like(backend.grad_views[0], 2.5)) + + +def test_failed_async_sync_closes_window_and_clears_pending_state(): + backend = _FakeAuthoritativeWrapper() + backend.cpu_infer.fail_next_sync = True + + backend.submit_backward_async(torch.ones(1, 1)) + with pytest.raises(RuntimeError, match=r"synthetic C\+\+ failure"): + backend.sync_backward() + + assert backend.parameters[0].grad is None + assert backend.validate_authoritative_optimizer_grad_state() == "closed" + assert backend._cache_depth == 0 + with pytest.raises(RuntimeError, match="No pending backward"): + backend.sync_backward() + + +def test_autograd_returns_no_base_gradient_and_preserves_published_alias(): + backend = _FakeAuthoritativeWrapper() + parameter = backend.parameters[0] + grad_view = backend.grad_views[0] + hidden_states = torch.ones(1, 1, 1, requires_grad=True) + expert_ids = torch.zeros(1, 1, dtype=torch.int64) + route_weights = torch.ones(1, 1, requires_grad=True) + + output = KTMoEFunction.apply( + hidden_states, + expert_ids, + route_weights, + backend, + parameter, + 1, + 1, + 0, + True, + False, + None, + False, + False, + parameter, + None, + None, + ) + output.sum().backward() + + # KTMoEFunction returned None for both references to the base Parameter; + # the alias published by the backend must therefore remain the sole grad. + assert parameter.grad is grad_view + torch.testing.assert_close(grad_view, torch.full_like(grad_view, 2.0)) From e26e21f0a263006161755c268b45f808428968ef Mon Sep 17 00:00:00 2001 From: yyj Date: Sat, 18 Jul 2026 23:06:35 +0800 Subject: [PATCH 20/20] [perf](kt-kernel): avoid eager Full-FT gradient zeroing Allocate authoritative Full-FT gradient buffers with torch.empty. The C++ state machine performs the mandatory full clear before first use, avoiding redundant Python-side first touch. --- kt-kernel/python/sft/base.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/kt-kernel/python/sft/base.py b/kt-kernel/python/sft/base.py index 9d14fd331..438867d74 100644 --- a/kt-kernel/python/sft/base.py +++ b/kt-kernel/python/sft/base.py @@ -420,10 +420,10 @@ def init_full_weight_grad_buffers( self.up_proj_buf = nn.Parameter(up_proj.to(dtype=dtype, device="cpu").contiguous(), requires_grad=True) self.down_proj_buf = nn.Parameter(down_proj.to(dtype=dtype, device="cpu").contiguous(), requires_grad=True) - # Create gradient buffers (C++ writes directly to these) - self.grad_gate_proj_buf = torch.zeros(E, I, H, dtype=dtype, device="cpu") - self.grad_up_proj_buf = torch.zeros(E, I, H, dtype=dtype, device="cpu") - self.grad_down_proj_buf = torch.zeros(E, H, I, dtype=dtype, device="cpu") + # C++ clears these authoritative gradient buffers before first use. + self.grad_gate_proj_buf = torch.empty(E, I, H, dtype=dtype, device="cpu") + self.grad_up_proj_buf = torch.empty(E, I, H, dtype=dtype, device="cpu") + self.grad_down_proj_buf = torch.empty(E, H, I, dtype=dtype, device="cpu") if self._uses_authoritative_optimizer_grads: self.register_authoritative_optimizer_grad("base.gate_proj", self.gate_proj_buf, self.grad_gate_proj_buf)