AD-Styles/mini-llava-demo
0
1"""MiniLLaVA — CLIP-ViT + MultiModalProjector + Qwen2.5 Causal LM.2 3LLaVA-1.5의 핵심 아키텍처를 직접 구현. HuggingFace의 LlavaForConditionalGeneration4같은 고수준 클래스를 사용하지 않고, 텍스트/이미지 임베딩 융합과 splice 로직을5저수준에서 직접 다룬다.6"""7from __future__ import annotations8 9import os10from typing import Optional11 12import torch13import torch.nn as nn14from transformers import (15 AutoModelForCausalLM,16 AutoTokenizer,17 CLIPImageProcessor,18 CLIPVisionModel,19)20 21from .config import IGNORE_INDEX, IMAGE_TOKEN, LLM_MODEL, VISION_MODEL22 23 24class MultiModalProjector(nn.Module):25 """CLIP의 시각 특징을 LLM의 임베딩 공간으로 매핑하는 2-layer MLP.26 27 LLaVA-1.5의 'mlp2x_gelu' projector를 그대로 따른다.28 """29 30 def __init__(self, vision_hidden_size: int, llm_hidden_size: int):31 super().__init__()32 self.fc1 = nn.Linear(vision_hidden_size, llm_hidden_size)33 self.act = nn.GELU()34 self.fc2 = nn.Linear(llm_hidden_size, llm_hidden_size)35 36 def forward(self, x: torch.Tensor) -> torch.Tensor:37 return self.fc2(self.act(self.fc1(x)))38 39 40class MiniLLaVA(nn.Module):41 """Vision-Language Model.42 43 - CLIP-ViT는 항상 frozen (강력한 사전학습 시각 표현 활용)44 - LLM은 기본 frozen (LLaVA Stage 1 alignment)45 - Projector만 학습 → 1.6M params 만으로 멀티모달 능력 부여46 """47 48 def __init__(49 self,50 vision_model_name: str = VISION_MODEL,51 llm_model_name: str = LLM_MODEL,52 freeze_vision: bool = True,53 freeze_llm: bool = True,54 torch_dtype: torch.dtype = torch.float32,55 ):56 super().__init__()57 58 self.vision = CLIPVisionModel.from_pretrained(vision_model_name)59 self.image_processor = CLIPImageProcessor.from_pretrained(vision_model_name)60 61 self.llm = AutoModelForCausalLM.from_pretrained(62 llm_model_name, torch_dtype=torch_dtype63 )64 self.tokenizer = AutoTokenizer.from_pretrained(llm_model_name)65 66 # <image> 플레이스홀더 추가67 if IMAGE_TOKEN not in self.tokenizer.get_vocab():68 self.tokenizer.add_special_tokens(69 {"additional_special_tokens": [IMAGE_TOKEN]}70 )71 self.llm.resize_token_embeddings(len(self.tokenizer))72 self.image_token_id = self.tokenizer.convert_tokens_to_ids(IMAGE_TOKEN)73 74 if self.tokenizer.pad_token_id is None:75 self.tokenizer.pad_token_id = self.tokenizer.eos_token_id76 77 vision_hidden = self.vision.config.hidden_size78 llm_hidden = self.llm.config.hidden_size79 self.projector = MultiModalProjector(vision_hidden, llm_hidden)80 81 if freeze_vision:82 for p in self.vision.parameters():83 p.requires_grad = False84 self.vision.eval()85 if freeze_llm:86 for p in self.llm.parameters():87 p.requires_grad = False88 89 # ──────────────────────────────────────────────────────────────────90 # Encoding91 # ──────────────────────────────────────────────────────────────────92 def encode_image(self, pixel_values: torch.Tensor) -> torch.Tensor:93 """[B, 3, H, W] → [B, N_patches, D_llm]. CLS 토큰 제외."""94 outputs = self.vision(pixel_values=pixel_values)95 patch_features = outputs.last_hidden_state[:, 1:, :]96 return self.projector(patch_features)97 98 # ──────────────────────────────────────────────────────────────────99 # Embedding fusion: <image> 위치를 patch tokens로 splice100 # ──────────────────────────────────────────────────────────────────101 def _merge(102 self,103 text_embeds: torch.Tensor,104 attention_mask: torch.Tensor,105 image_embeds: torch.Tensor,106 input_ids: torch.Tensor,107 labels: Optional[torch.Tensor] = None,108 ):109 """input_ids에서 <image> 위치를 image_embeds(N개 patch)로 교체.110 111 - 모든 샘플은 정확히 1개의 <image> 토큰을 가진다고 가정112 - text/mask/label을 모두 일관되게 재정렬113 """114 B, L, D = text_embeds.shape115 N = image_embeds.shape[1]116 new_L = L - 1 + N117 118 device = text_embeds.device119 merged_embeds = torch.zeros(B, new_L, D, dtype=text_embeds.dtype, device=device)120 merged_mask = torch.zeros(B, new_L, dtype=attention_mask.dtype, device=device)121 merged_labels = (122 torch.full((B, new_L), IGNORE_INDEX, dtype=torch.long, device=device)123 if labels is not None124 else None125 )126 127 for b in range(B):128 img_pos = (input_ids[b] == self.image_token_id).nonzero(as_tuple=True)[0]129 if len(img_pos) != 1:130 raise ValueError(131 f"sample {b}는 <image> 토큰이 {len(img_pos)}개 — 정확히 1개여야 합니다."132 )133 p = img_pos.item()134 135 # 앞 / 이미지 / 뒤 순으로 splice136 merged_embeds[b, :p] = text_embeds[b, :p]137 merged_embeds[b, p : p + N] = image_embeds[b]138 merged_embeds[b, p + N :] = text_embeds[b, p + 1 :]139 140 merged_mask[b, :p] = attention_mask[b, :p]141 merged_mask[b, p : p + N] = 1142 merged_mask[b, p + N :] = attention_mask[b, p + 1 :]143 144 if labels is not None:145 merged_labels[b, :p] = labels[b, :p]146 # 이미지 patch 위치는 IGNORE_INDEX 유지 (이미 채워둠)147 merged_labels[b, p + N :] = labels[b, p + 1 :]148 149 return merged_embeds, merged_mask, merged_labels150 151 # ──────────────────────────────────────────────────────────────────152 # Forward (학습)153 # ──────────────────────────────────────────────────────────────────154 def forward(155 self,156 input_ids: torch.Tensor,157 attention_mask: torch.Tensor,158 pixel_values: torch.Tensor,159 labels: Optional[torch.Tensor] = None,160 ):161 text_embeds = self.llm.get_input_embeddings()(input_ids)162 image_embeds = self.encode_image(pixel_values)163 164 merged_embeds, merged_mask, merged_labels = self._merge(165 text_embeds, attention_mask, image_embeds, input_ids, labels166 )167 168 return self.llm(169 inputs_embeds=merged_embeds,170 attention_mask=merged_mask,171 labels=merged_labels,172 return_dict=True,173 )174 175 # ──────────────────────────────────────────────────────────────────176 # Generation (추론)177 # ──────────────────────────────────────────────────────────────────178 @torch.no_grad()179 def generate(180 self,181 input_ids: torch.Tensor,182 attention_mask: torch.Tensor,183 pixel_values: torch.Tensor,184 max_new_tokens: int = 128,185 temperature: float = 0.7,186 top_p: float = 0.9,187 do_sample: bool = True,188 ) -> torch.Tensor:189 text_embeds = self.llm.get_input_embeddings()(input_ids)190 image_embeds = self.encode_image(pixel_values)191 merged_embeds, merged_mask, _ = self._merge(192 text_embeds, attention_mask, image_embeds, input_ids, labels=None193 )194 195 return self.llm.generate(196 inputs_embeds=merged_embeds,197 attention_mask=merged_mask,198 max_new_tokens=max_new_tokens,199 temperature=temperature,200 top_p=top_p,201 do_sample=do_sample,202 pad_token_id=self.tokenizer.pad_token_id,203 eos_token_id=self.tokenizer.eos_token_id,204 )205 206 # ──────────────────────────────────────────────────────────────────207 # Checkpoint I/O — projector만 저장 (LLM/CLIP은 HF에서 다시 로드)208 # ──────────────────────────────────────────────────────────────────209 def save_projector(self, path: str) -> None:210 os.makedirs(os.path.dirname(path), exist_ok=True)211 torch.save(self.projector.state_dict(), path)212 213 def load_projector(self, path: str, map_location: str = "cpu") -> None:214 state = torch.load(path, map_location=map_location)215 self.projector.load_state_dict(state)216 217 def load_lora_adapter(self, adapter_path: str) -> None:218 """학습된 LoRA adapter를 frozen LLM 위에 부착."""219 from peft import PeftModel220 221 self.llm = PeftModel.from_pretrained(self.llm, adapter_path)222 self.llm.eval()223 224 def trainable_parameters(self):225 return [p for p in self.parameters() if p.requires_grad]226 227 def num_trainable(self) -> int:228 return sum(p.numel() for p in self.trainable_parameters())229 