baguettotron-internvit-alignment / modeling_baguettotron_vlm.py
andreagemelli's picture
Fix generation from the published artifact; rewrite model card
2c4010c verified
Raw
History Blame
12.3 kB
"""BaguettotronVLM model — self-contained for HuggingFace Hub."""
from __future__ import annotations
import torch
import torch.nn as nn
from transformers import (
AutoModel,
AutoModelForCausalLM,
AutoTokenizer,
LogitsProcessor,
PreTrainedModel,
StoppingCriteria,
)
from transformers.generation.logits_process import RepetitionPenaltyLogitsProcessor
from transformers.modeling_outputs import CausalLMOutputWithPast
from .configuration_baguettotron_vlm import BaguettotronVLMConfig
class SafeRepetitionPenalty(LogitsProcessor):
"""Repetition penalty that is a no-op on the first decoding step.
generate(inputs_embeds=...) begins with an empty input_ids tensor. On MPS
the empty-index gather inside transformers' RepetitionPenaltyLogitsProcessor
zeroes the entire logits row rather than leaving it untouched (it is a
correct no-op on CPU and CUDA), which corrupts the first token and derails
the answer. There is nothing to penalise on that step anyway, so skipping it
removes the corruption while keeping the penalty for every later step.
"""
def __init__(self, penalty: float):
self.inner = RepetitionPenaltyLogitsProcessor(penalty)
def __call__(self, input_ids: torch.Tensor, scores: torch.Tensor) -> torch.Tensor:
if input_ids.shape[-1] == 0:
return scores
return self.inner(input_ids, scores)
class StopOnTurnEnd(StoppingCriteria):
"""Stop when the decoded tail contains a turn marker.
These weights never emit a stop *token*: <|im_end|> was masked out of the
training loss, so the model learned to spell the marker out as ordinary
text ("<|", "im", "_", "end", "|>") and then start a fresh turn. Passing
eos_token_id therefore does nothing — generation runs to max_new_tokens
every time, which is ~20x more tokens than the answer needs.
Matching on decoded text rather than token ids is deliberate: how the
marker splits depends on what precedes it (".<|" merges into a single
token after a period), so an id-sequence match misses many cases.
"""
def __init__(self, tokenizer, markers=("<|im_end|>", "<|im_start|>"), window=8):
self.tokenizer = tokenizer
self.markers = markers
self.window = window
def __call__(self, input_ids: torch.Tensor, scores, **kwargs) -> bool:
tail = self.tokenizer.decode(
input_ids[0, -self.window:], skip_special_tokens=False
)
return any(marker in tail for marker in self.markers)
class PixelUnshuffleProjector(nn.Module):
"""Reduces ViT tokens 4× via PixelUnshuffle then projects to LLM dim."""
def __init__(self, in_dim: int, out_dim: int, factor: int):
super().__init__()
self.factor = factor
self.unshuffle = nn.PixelUnshuffle(factor)
self.mlp = nn.Sequential(
nn.Linear(in_dim * factor * factor, out_dim),
nn.GELU(),
nn.Linear(out_dim, out_dim),
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
B, N, D = x.shape
spatial = int(N ** 0.5)
x = x.reshape(B, spatial, spatial, D).permute(0, 3, 1, 2)
x = self.unshuffle(x)
x = x.flatten(2).transpose(1, 2)
return self.mlp(x)
class BaguettotronVLMForConditionalGeneration(PreTrainedModel):
"""
BaguettotronVLM: InternViT-300M + PixelUnshuffle projector + Baguettotron-321M.
~628M total parameters. The same architecture is shipped for every
checkpoint; only the weights and `config.chat_style` differ.
Load with:
from transformers import AutoModelForImageTextToText
model = AutoModelForImageTextToText.from_pretrained(
"andreagemelli/baguettotron-vision-vqa",
trust_remote_code=True,
dtype=torch.bfloat16,
)
"""
config_class = BaguettotronVLMConfig
_no_split_modules = ["InternVisionEncoderLayer", "LlamaDecoderLayer"]
# Nothing is tied in this model — see __init__. Without this, the inherited
# Llama tied-weight bookkeeping drops llm.lm_head.weight during loading.
_tied_weights_keys: list[str] = []
# Tell HF Trainer not to pass num_items_in_batch (loss handled internally)
model_accepts_loss_kwargs: bool = False
def __init__(self, config: BaguettotronVLMConfig):
super().__init__(config)
# Honour the dtype the caller asked for. from_pretrained(dtype=X) makes X
# the default dtype for the duration of __init__, and the projector picks
# it up automatically. Hard-coding bfloat16 for the submodules while the
# projector follows the default is what made dtype=torch.float32 fail with
# "mat1 and mat2 must have the same dtype".
dtype = torch.get_default_dtype()
if dtype not in (torch.float32, torch.float16, torch.bfloat16):
dtype = torch.bfloat16
# The explicit CPU device context shields these nested from_pretrained
# calls from an outer meta-device init context (accelerate's device_map,
# and transformers >= 5), which otherwise aborts with "You are using
# from_pretrained with a meta device context manager".
with torch.device("cpu"):
self.vit = AutoModel.from_pretrained(
config.vit_model_id,
dtype=dtype,
low_cpu_mem_usage=True,
trust_remote_code=True,
)
self.projector = PixelUnshuffleProjector(
in_dim=config.vit_hidden,
out_dim=config.llm_hidden,
factor=config.unshuffle_factor,
)
tokenizer = AutoTokenizer.from_pretrained(config.llm_model_id)
tokenizer.add_special_tokens(
{"additional_special_tokens": [config.image_token, "</image>"]}
)
raw_id = tokenizer.convert_tokens_to_ids(config.image_token)
self.image_token_id: int = (
raw_id if isinstance(raw_id, int) else int(raw_id[0])
)
self.llm = AutoModelForCausalLM.from_pretrained(
config.llm_model_id, dtype=dtype
)
self.llm.resize_token_embeddings(len(tokenizer), mean_resizing=False)
# Break weight tying — safetensors rejects shared-storage tensors, and
# training updates lm_head independently of the input embeddings.
self.llm.lm_head.weight = nn.Parameter(self.llm.lm_head.weight.data.clone())
# Keep it broken. Llama declares lm_head.weight as a tied key, so
# from_pretrained would skip it while loading and then re-tie it to
# embed_tokens — silently discarding the trained output head.
self.llm.config.tie_word_embeddings = False
self.llm._tied_weights_keys = []
# NOTE: transformers >= 5 cannot load this model, and the blocker is
# upstream: OpenGVLab's InternViT remote code predates v5's tied-weight
# API, so v5 aborts inside the nested AutoModel.from_pretrained above
# with "'InternVisionModel' object has no attribute
# 'all_tied_weights_keys'". Nothing can be patched from here — the model
# cards pin transformers<5 until InternViT is updated.
self._tokenizer = tokenizer
def _init_weights(self, module: nn.Module) -> None:
# Pretrained components are initialised from their respective hubs;
# the projector weights come from the saved checkpoint — skip random init.
pass
# ------------------------------------------------------------------
# Training interface
# ------------------------------------------------------------------
def forward(
self,
input_ids: torch.Tensor,
attention_mask: torch.Tensor,
labels: torch.Tensor | None = None,
pixel_values: torch.Tensor | None = None,
**kwargs,
) -> CausalLMOutputWithPast:
inputs_embeds = self.llm.get_input_embeddings()(input_ids)
if pixel_values is not None:
pixel_values = pixel_values.to(dtype=self.vit.dtype, device=self.vit.device)
with torch.no_grad():
vit_out = self.vit(pixel_values)
image_features = vit_out.last_hidden_state
if image_features.shape[1] == self.config.vit_tokens + 1:
image_features = image_features[:, 1:, :]
visual_tokens = self.projector(image_features.to(inputs_embeds.dtype))
image_mask = input_ids == self.image_token_id
inputs_embeds[image_mask] = visual_tokens.reshape(
-1, self.config.llm_hidden
).to(inputs_embeds.dtype)
return self.llm(
inputs_embeds=inputs_embeds,
attention_mask=attention_mask,
labels=labels,
return_dict=True,
use_cache=False,
)
# ------------------------------------------------------------------
# Inference interface
# ------------------------------------------------------------------
@torch.no_grad()
def chat(
self,
input_ids: torch.Tensor,
attention_mask: torch.Tensor,
pixel_values: torch.Tensor | None = None,
max_new_tokens: int = 256,
repetition_penalty: float = 1.3,
**generate_kwargs,
) -> str:
"""Inject visual tokens, generate autoregressively, return decoded string."""
inputs_embeds = self.llm.get_input_embeddings()(input_ids)
if pixel_values is not None:
pixel_values = pixel_values.to(dtype=self.vit.dtype, device=self.vit.device)
vit_out = self.vit(pixel_values)
image_features = vit_out.last_hidden_state
if image_features.shape[1] == self.config.vit_tokens + 1:
image_features = image_features[:, 1:, :]
visual_tokens = self.projector(image_features.to(inputs_embeds.dtype))
image_mask = input_ids == self.image_token_id
inputs_embeds[image_mask] = visual_tokens.reshape(
-1, self.config.llm_hidden
).to(inputs_embeds.dtype)
# These weights end a turn by spelling out the literal text "<|im_end|>"
# as ordinary tokens and then emitting <|end_of_text|> — the special
# <|im_end|> token is never produced. <|end_of_text|> is therefore the id
# that actually stops generation; the others are listed for safety.
# PleIAs/Baguettotron's tokenizer declares no eos_token, so eos_token_id
# is None there and must be filtered out — an unfiltered [None, ...] makes
# generate() raise "'NoneType' object cannot be interpreted as an integer".
unk_id = self._tokenizer.unk_token_id
candidates = (
self._tokenizer.eos_token_id,
self._tokenizer.convert_tokens_to_ids("<|end_of_text|>"),
self._tokenizer.convert_tokens_to_ids("<|im_end|>"),
)
eos_ids = list(
dict.fromkeys(
tok_id
for tok_id in candidates
if tok_id is not None and tok_id != unk_id
)
)
# do_sample defaults to greedy but can be overridden by callers
# (e.g. the inference sweep) without colliding on the keyword.
generate_kwargs.setdefault("do_sample", False)
generate_kwargs.setdefault(
"stopping_criteria", [StopOnTurnEnd(self._tokenizer)]
)
# The penalty goes through SafeRepetitionPenalty rather than generate()'s
# repetition_penalty kwarg, so the first decoding step is skipped.
if repetition_penalty != 1.0:
generate_kwargs.setdefault(
"logits_processor", [SafeRepetitionPenalty(repetition_penalty)]
)
output_ids = self.llm.generate(
inputs_embeds=inputs_embeds,
attention_mask=attention_mask,
max_new_tokens=max_new_tokens,
eos_token_id=eos_ids,
**generate_kwargs,
)
decoded = self._tokenizer.decode(output_ids[0], skip_special_tokens=False)
for marker in ("<|im_end|>", "<|im_start|>", "<|end_of_text|>"):
if marker in decoded:
decoded = decoded[: decoded.index(marker)]
return decoded.strip()