Skip to content

[NPU] Add Kimi-K2-Thinking support #395

Description

@yuxinshan

Your Question

[NPU] To support Kimi-K2-Thinking model.

Traceback (most recent call last):
  File "/tmp/ray_vime_npu_kimi_k2/session_2026-xx-xx_02-17-43_421881_2412330/runtime_resources/working_dir_files/_ray_pkg_d6f271e2f85b1e48/train.py", line 111, in <module>
    train(args)
  File "/tmp/ray_vime_npu_kimi_k2/session_2026-xx-xx_02-17-43_421881_2412330/runtime_resources/working_dir_files/_ray_pkg_d6f271e2f85b1e48/train.py", line 28, in train
    actor_model, critic_model = create_training_models(args, pgs, rollout_manager)
                                ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "/tmp/ray_vime_npu_kimi_k2/session_2026-xx-xx_02-17-43_421881_2412330/runtime_resources/working_dir_files/_ray_pkg_d6f271e2f85b1e48/vime/ray/placement_group.py", line 182, in create_training_models
    actor_start_rollout_ids = ray.get(
                              ^^^^^^^^
  File "/usr/local/python3.12.13/lib/python3.12/site-packages/ray/_private/auto_init_hook.py", line 22, in auto_init_wrapper
    return fn(*args, **kwargs)
           ^^^^^^^^^^^^^^^^^^^
  File "/usr/local/python3.12.13/lib/python3.12/site-packages/ray/_private/client_mode_hook.py", line 104, in wrapper
    return func(*args, **kwargs)
           ^^^^^^^^^^^^^^^^^^^^^
  File "/usr/local/python3.12.13/lib/python3.12/site-packages/ray/_private/worker.py", line 2858, in get
    values, debugger_breakpoint = worker.get_objects(object_refs, timeout=timeout)
                                  ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "/usr/local/python3.12.13/lib/python3.12/site-packages/ray/_private/worker.py", line 958, in get_objects
    raise value.as_instanceof_cause()
ray.exceptions.RayTaskError(AttributeError): ray::MegatronTrainRayActor.init() (pid=2438544, ip=xx.xx.xx.xx, actor_id=9408962f7b41bfe7cc8ec0ae03000000, repr=<vime.backends.megatron_utils.actor.MegatronTrainRayActor object at 0xffcfd2351ee0>)
           ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
           ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "/tmp/ray_vime_npu_kimi_k2/session_2026-xx-xx_02-17-43_421881_2412330/runtime_resources/working_dir_files/_ray_pkg_d6f271e2f85b1e48/vime/utils/timer.py", line 97, in wrapper
    return fn(*args, **kwargs)
           ^^^^^^^^^^^^^^^^^^^
  File "/tmp/ray_vime_npu_kimi_k2/session_2026-xx-xx_02-17-43_421881_2412330/runtime_resources/working_dir_files/_ray_pkg_d6f271e2f85b1e48/vime/backends/megatron_utils/actor.py", line 107, in init
    self.model, self.optimizer, self.opt_param_scheduler, loaded_rollout_id = initialize_model_and_optimizer(
                                                                              ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "/tmp/ray_vime_npu_kimi_k2/session_2026-xx-xx_02-17-43_421881_2412330/runtime_resources/working_dir_files/_ray_pkg_d6f271e2f85b1e48/vime/backends/megatron_utils/model.py", line 946, in initialize_model_and_optimizer
    iteration, _ = load_checkpoint(
                   ^^^^^^^^^^^^^^^^
  File "/tmp/ray_vime_npu_kimi_k2/session_2026-xx-xx_02-17-43_421881_2412330/runtime_resources/working_dir_files/_ray_pkg_d6f271e2f85b1e48/vime/backends/megatron_utils/checkpoint.py", line 127, in load_checkpoint
    return _load_checkpoint_hf(
           ^^^^^^^^^^^^^^^^^^^^
  File "/tmp/ray_vime_npu_kimi_k2/session_2026-xx-xx_02-17-43_421881_2412330/runtime_resources/working_dir_files/_ray_pkg_d6f271e2f85b1e48/vime/backends/megatron_utils/checkpoint.py", line 153, in _load_checkpoint_hf
    bridge.load_hf_weights(ddp_model)
  File "/root/Megatron-Bridge/src/megatron/bridge/models/conversion/auto_bridge.py", line 340, in load_hf_weights
    self._model_bridge.load_weights_hf_to_megatron(
  File "/root/Megatron-Bridge/src/megatron/bridge/models/conversion/model_bridge.py", line 795, in load_weights_hf_to_megatron
    converted_weights = task.mapping.hf_to_megatron(hf_weights, task.megatron_module)
                        ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "/root/Megatron-Bridge/src/megatron/bridge/models/conversion/param_mapping.py", line 1246, in hf_to_megatron
    self._detected_type = self._detect_parallelism_type(megatron_module)
                          ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "/root/Megatron-Bridge/src/megatron/bridge/models/conversion/param_mapping.py", line 1215, in _detect_parallelism_type
    if module.parallel_mode == "column":
       ^^^^^^^^^^^^^^^^^^^^
  File "/usr/local/python3.12.13/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1965, in __getattr__
    raise AttributeError(
AttributeError: 'MindSpeedTELinear' object has no attribute 'parallel_mode'

What I've Tried

Try to run the Kimi-K2-Thinking model:
#399

CKPT_ARGS=(   
   --hf-checkpoint $BASE_DIR/Kimi-K2-Thinking/
   --ref-load $BASE_DIR/Kimi-K2-Thinking/
   --no-load-optim
)

ROLLOUT_ARGS=(
   --prompt-data $BASE_DIR/dapo-math-17k/dapo-math-17k.jsonl
   --input-key prompt
   --label-key label
   --apply-chat-template
   --rollout-shuffle
   --rm-type math
   --num-rollout 100
   --rollout-batch-size 128
   --n-samples-per-prompt 8
   --rollout-max-response-len 16384
   --rollout-temperature 1
   --over-sampling-batch-size 256
   --num-steps-per-rollout 4
)

PERF_ARGS=(
   --tensor-model-parallel-size 8
   --sequence-parallel
   --pipeline-model-parallel-size 8
   --context-parallel-size 4
   --expert-model-parallel-size 32
   --expert-tensor-parallel-size 1
   --use-dynamic-batch-size
   --max-tokens-per-gpu 16384
)

GRPO_ARGS=(
   --advantage-estimator grpo
   --use-kl-loss
   --kl-loss-coef 0.00
   --kl-loss-type low_var_kl
   --entropy-coef 0.00
   --eps-clip 0.2
   --eps-clip-high 0.28
)

OPTIMIZER_ARGS=(
   --optimizer adam
   --lr 1e-6
   --lr-decay-style constant
   --weight-decay 0.1
   --adam-beta1 0.9
   --adam-beta2 0.98
   --optimizer-cpu-offload
   --overlap-cpu-optimizer-d2h-h2d
   --use-precision-aware-optimizer
)

VLLM_ARGS=(
   --rollout-backend vllm
   --rollout-num-gpus-per-engine 8
   --vllm-gpu-memory-utilization 0.45
   --vllm-weight-sync-mode native
   --vllm-enable-sleep-mode
   --vllm-tensor-parallel-size 8
   --vllm-data-parallel-size 8
   --vllm-enable-expert-parallel
   --vllm-server-concurrency 1024
   --vllm-enable-cuda-graph
   --vllm-quantization None
)


MISC_ARGS=(
   --attention-dropout 0.0
   --hidden-dropout 0.0
   --accumulate-allreduce-grads-in-fp32
   --attention-softmax-in-fp32
   --attention-backend flash
   --micro-batch-size 1
   --use-flash-attn
   --gradient-accumulation-steps 8
   --no-gradient-accumulation-fusion
   --train-memory-margin-bytes 8589934592
   --megatron-to-hf-mode bridge
)

Environment (if relevant)

  • vime version: ascend branch
  • Python version: 3.12
  • PyTorch version: 2.10.0
  • vLLM version: 0.22.1
  • vLLM-Ascend version: 0.22.1.rc1
  • Megatron version: 0.16.0rc0
  • Mindspeeed version: 0.14.1

Additional Context

No response

Pre-submission Checklist

Metadata

Metadata

Assignees

No one assigned

    Labels

    questionFurther information is requested

    Type

    No type

    Projects

    No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions