Enable Kimi-K2.5 piecewise CUDA graph (#26382)

This commit is contained in:
Xiaoyu Zhang
2026-05-27 22:51:33 -07:00
committed by GitHub
parent 5018e5c969
commit e60f799b40
2 changed files with 14 additions and 0 deletions
+1
View File
@@ -178,6 +178,7 @@ def _forward_with_allreduce_fusion(
residual=residual,
weight=weight,
eps=norm_module.variance_epsilon,
max_token_num=max(x.shape[0], 2048),
use_attn_tp_group=use_attn_tp_group,
)
if fused_result[0] is not None:
+13
View File
@@ -674,6 +674,19 @@ class KimiK25ForConditionalGeneration(nn.Module):
self.vision_tower = self.vision_tower.to(dtype=target_dtype)
self.mm_projector = self.mm_projector.to(dtype=target_dtype)
@property
def model(self):
# Alias .model to .language_model so this class satisfies the piecewise
# CUDA graph gate, which checks `hasattr(model, "model")`.
return self.language_model
def __setattr__(self, name, value):
# Skip redundant self.model.model assignment in runner to avoid duplicate
# nn.Module registration.
if name == "model":
return
super().__setattr__(name, value)
def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
device = self.vision_tower.device
target_dtype = self.vision_tower.patch_embed.proj.weight.dtype