diff --git a/python/sglang/srt/models/kimi_k25.py b/python/sglang/srt/models/kimi_k25.py index f1b38b1fe..4b5dbbdfc 100644 --- a/python/sglang/srt/models/kimi_k25.py +++ b/python/sglang/srt/models/kimi_k25.py @@ -786,5 +786,34 @@ class KimiK25ForConditionalGeneration(nn.Module): num_groups=text_config.n_group, ) + def set_eagle3_layers_to_capture( + self, layer_ids: Optional[List[int]] = None + ) -> None: + """Set the layers to capture for EAGLE3 speculative decoding.""" + if not hasattr(self.language_model, "set_eagle3_layers_to_capture"): + raise AttributeError( + "language_model does not support EAGLE3 speculative decoding." + ) + + self.language_model.set_eagle3_layers_to_capture(layer_ids) + + def get_embed_and_head(self) -> Tuple[torch.Tensor, torch.Tensor]: + """Get embedding and LM head weights for speculative decoding.""" + if not hasattr(self.language_model, "get_embed_and_head"): + raise AttributeError( + "language_model does not support get_embed_and_head()." + ) + + return self.language_model.get_embed_and_head() + + def set_embed_and_head(self, embed: torch.Tensor, head: torch.Tensor) -> None: + """Set embedding and LM head weights for speculative decoding.""" + if not hasattr(self.language_model, "set_embed_and_head"): + raise AttributeError( + "language_model does not support set_embed_and_head()." + ) + + self.language_model.set_embed_and_head(embed, head) + EntryClass = [KimiK25ForConditionalGeneration]