hugging-apps/echo-memory
0
1from transformers import LlamaModel, LlamaConfig, DynamicCache, LlavaForConditionalGeneration2from copy import deepcopy3import torch4 5 6class HunyuanVideoLLMEncoder(LlamaModel):7 8 def __init__(self, config: LlamaConfig):9 super().__init__(config)10 self.auto_offload = False11 12 def enable_auto_offload(self, **kwargs):13 self.auto_offload = True14 15 def forward(self, input_ids, attention_mask, hidden_state_skip_layer=2):16 embed_tokens = deepcopy(self.embed_tokens).to(input_ids.device) if self.auto_offload else self.embed_tokens17 inputs_embeds = embed_tokens(input_ids)18 19 past_key_values = DynamicCache()20 21 cache_position = torch.arange(0, inputs_embeds.shape[1], device=inputs_embeds.device)22 position_ids = cache_position.unsqueeze(0)23 24 causal_mask = self._update_causal_mask(attention_mask, inputs_embeds, cache_position, None, False)25 hidden_states = inputs_embeds26 27 # create position embeddings to be shared across the decoder layers28 rotary_emb = deepcopy(self.rotary_emb).to(input_ids.device) if self.auto_offload else self.rotary_emb29 position_embeddings = rotary_emb(hidden_states, position_ids)30 31 # decoder layers32 for layer_id, decoder_layer in enumerate(self.layers):33 if self.auto_offload:34 decoder_layer = deepcopy(decoder_layer).to(hidden_states.device)35 layer_outputs = decoder_layer(36 hidden_states,37 attention_mask=causal_mask,38 position_ids=position_ids,39 past_key_value=past_key_values,40 output_attentions=False,41 use_cache=True,42 cache_position=cache_position,43 position_embeddings=position_embeddings,44 )45 hidden_states = layer_outputs[0]46 if layer_id + hidden_state_skip_layer + 1 >= len(self.layers):47 break48 49 return hidden_states50 51 52class HunyuanVideoMLLMEncoder(LlavaForConditionalGeneration):53 54 def __init__(self, config):55 super().__init__(config)56 self.auto_offload = False57 58 def enable_auto_offload(self, **kwargs):59 self.auto_offload = True60 61 # TODO: implement the low VRAM inference for MLLM.62 def forward(self, input_ids, pixel_values, attention_mask, hidden_state_skip_layer=2):63 outputs = super().forward(input_ids=input_ids,64 attention_mask=attention_mask,65 output_hidden_states=True,66 pixel_values=pixel_values)67 hidden_state = outputs.hidden_states[-(hidden_state_skip_layer + 1)]68 return hidden_state69 