Nexa-AI-VL-4x4B-Base / modeling_nexa.py
Neura-Tech-AI's picture
Update modeling_nexa.py
ea8ac28 verified
Raw
History Blame Contribute Delete
4.23 kB
import os
import torch
import torch.nn as nn
from transformers import PreTrainedModel, AutoModel, AutoModelForCausalLM
from safetensors.torch import load_file
from .configuration_nexa import NexaVisionMoEConfig
class MultimodalProjector(nn.Module):
def __init__(self, config):
super().__init__()
self.proj = nn.Sequential(
nn.Linear(config.vision_dim, config.llm_dim),
nn.GELU(),
nn.Linear(config.llm_dim, config.llm_dim)
)
def forward(self, x):
return self.proj(x)
class NexaVisionMoEForConditionalGeneration(PreTrainedModel):
config_class = NexaVisionMoEConfig
def __init__(self, config):
super().__init__(config)
self.config = config
# Base Vision Tower
self.vision_tower = AutoModel.from_pretrained(
config.vision_model_id,
torch_dtype=torch.float16
).vision_model
# Projector
self.projector = MultimodalProjector(config).to(torch.float16)
# LLM Sub-model (Loads local base config)
self.language_model = AutoModelForCausalLM.from_pretrained(
config.llm_model_id,
torch_dtype=torch.float16
)
# Updated weights loading logic with renamed file: vision_projector.safetensors
current_dir = os.path.dirname(__file__)
vis_proj_path = os.path.join(current_dir, "vision_projector.safetensors")
if os.path.exists(vis_proj_path):
state_dict = load_file(vis_proj_path)
vis_state = {k.replace("vision_tower.", ""): v for k, v in state_dict.items() if k.startswith("vision_tower.")}
proj_state = {k.replace("projector.", ""): v for k, v in state_dict.items() if k.startswith("projector.")}
if vis_state:
self.vision_tower.load_state_dict(vis_state, strict=False)
if proj_state:
self.projector.load_state_dict(proj_state, strict=False)
def get_input_embeddings(self):
return self.language_model.get_input_embeddings()
def set_input_embeddings(self, value):
self.language_model.set_input_embeddings(value)
def forward(self, input_ids=None, pixel_values=None, attention_mask=None, labels=None, **kwargs):
inputs_embeds = self.get_input_embeddings()(input_ids)
if pixel_values is not None:
vision_outputs = self.vision_tower(pixel_values).last_hidden_state
image_embeds = self.projector(vision_outputs.to(inputs_embeds.dtype))
inputs_embeds = torch.cat([image_embeds, inputs_embeds], dim=1)
if attention_mask is not None:
img_attn = torch.ones((attention_mask.shape[0], image_embeds.shape[1]), device=attention_mask.device, dtype=attention_mask.dtype)
attention_mask = torch.cat([img_attn, attention_mask], dim=1)
return self.language_model(
inputs_embeds=inputs_embeds,
attention_mask=attention_mask,
labels=labels,
**kwargs
)
def generate(self, pixel_values=None, input_ids=None, attention_mask=None, **kwargs):
if pixel_values is not None:
vision_outputs = self.vision_tower(pixel_values).last_hidden_state
image_embeds = self.projector(vision_outputs.to(self.dtype))
text_embeds = self.get_input_embeddings()(input_ids)
inputs_embeds = torch.cat([image_embeds, text_embeds], dim=1)
if attention_mask is not None:
img_attn = torch.ones((attention_mask.shape[0], image_embeds.shape[1]), device=attention_mask.device, dtype=attention_mask.dtype)
attention_mask = torch.cat([img_attn, attention_mask], dim=1)
return self.language_model.generate(
inputs_embeds=inputs_embeds,
attention_mask=attention_mask,
**kwargs
)
return self.language_model.generate(
input_ids=input_ids,
attention_mask=attention_mask,
**kwargs
)