CoolFace
Apppublic

AD-Styles/mini-llava-demo

sourceHugging Facemitupdated 5mo agoView on Hugging Face
0likes
model.py229 linesDownload Raw Back to src
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