[LFM2-VL] Add DSpark speculative decoding (#40651)
This commit is contained in:
@@ -180,6 +180,14 @@ class Lfm2VlForConditionalGeneration(nn.Module):
|
|||||||
def get_input_embeddings(self) -> nn.Embedding:
|
def get_input_embeddings(self) -> nn.Embedding:
|
||||||
return self.language_model.model.embed_tokens
|
return self.language_model.model.embed_tokens
|
||||||
|
|
||||||
|
@property
|
||||||
|
def lm_head(self):
|
||||||
|
return self.language_model.lm_head
|
||||||
|
|
||||||
|
def set_dflash_layers_to_capture(self, layer_ids: List[int]):
|
||||||
|
# Lfm2ForCausalLM applies the HF-layer-k -> "before layer k+1" shift.
|
||||||
|
self.language_model.set_dflash_layers_to_capture(layer_ids)
|
||||||
|
|
||||||
def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
|
def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
|
||||||
"""Process images through vision tower and projector.
|
"""Process images through vision tower and projector.
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user