cbensimon/moondream1
0
1from __future__ import annotations2 3import spaces4import torch5from PIL import Image6from einops import rearrange7from torchvision.transforms.v2 import (8 Compose,9 Resize,10 InterpolationMode,11 ToImage,12 ToDtype,13 Normalize,14)15 16from transformers import CodeGenTokenizerFast as Tokenizer17from accelerate import init_empty_weights, load_checkpoint_and_dispatch18import re19 20import math21from typing import Optional22 23from transformers import PretrainedConfig24 25 26import math27from dataclasses import dataclass, field28from typing import Any, Dict, Optional, Tuple, Union29 30import torch31import torch.nn as nn32from einops import rearrange, repeat33from transformers import PretrainedConfig, PreTrainedModel34from transformers.activations import ACT2FN35from transformers.modeling_outputs import CausalLMOutputWithPast36 37pad_input, unpad_input = None, None38FlashRotaryEmbedding = None39FlashSelfAttention, FlashCrossAttention = None, None40FusedDense = None41 42if torch.cuda.is_available():43 DEVICE = "cuda"44 DTYPE = torch.float1645else:46 DEVICE = "cpu"47 DTYPE = torch.float3248 49 50class PhiConfig(PretrainedConfig):51 """Phi configuration."""52 53 model_type = "phi-msft"54 attribute_map = {55 "max_position_embeddings": "n_positions",56 "hidden_size": "n_embd",57 "num_attention_heads": "n_head",58 "num_hidden_layers": "n_layer",59 }60 61 def __init__(62 self,63 vocab_size: int = 50304,64 n_positions: int = 2048,65 n_embd: int = 1024,66 n_layer: int = 20,67 n_inner: Optional[int] = None,68 n_head: int = 16,69 n_head_kv: Optional[int] = None,70 rotary_dim: Optional[int] = 32,71 activation_function: Optional[str] = "gelu_new",72 flash_attn: bool = False,73 flash_rotary: bool = False,74 fused_dense: bool = False,75 attn_pdrop: float = 0.0,76 embd_pdrop: float = 0.0,77 resid_pdrop: float = 0.0,78 layer_norm_epsilon: float = 1e-5,79 initializer_range: float = 0.02,80 tie_word_embeddings: bool = False,81 pad_vocab_size_multiple: int = 64,82 gradient_checkpointing: bool = False,83 **kwargs,84 ) -> None:85 self.vocab_size = int(86 math.ceil(vocab_size / pad_vocab_size_multiple) * pad_vocab_size_multiple87 )88 self.n_positions = n_positions89 self.n_embd = n_embd90 self.n_layer = n_layer91 self.n_inner = n_inner92 self.n_head = n_head93 self.n_head_kv = n_head_kv94 self.rotary_dim = min(rotary_dim, n_embd // n_head)95 self.activation_function = activation_function96 self.flash_attn = flash_attn97 self.flash_rotary = flash_rotary98 self.fused_dense = fused_dense99 self.attn_pdrop = attn_pdrop100 self.embd_pdrop = embd_pdrop101 self.resid_pdrop = resid_pdrop102 self.layer_norm_epsilon = layer_norm_epsilon103 self.initializer_range = initializer_range104 self.gradient_checkpointing = gradient_checkpointing105 106 super().__init__(tie_word_embeddings=tie_word_embeddings, **kwargs)107 108 109@dataclass110class InferenceParams:111 """Inference parameters passed to model to efficiently calculate112 and store context during inference.113 114 Reference:115 https://github.com/Dao-AILab/flash-attention/blob/main/flash_attn/utils/generation.py.116 117 Args:118 max_seqlen: Maximum sequence length.119 max_batch_size: Maximum batch size.120 seqlen_offset: Sequence length offset.121 batch_size_offset: Batch size offset.122 key_value_memory_dict: Key value memory dictionary.123 lengths_per_sample: Lengths per sample.124 125 """126 127 max_seqlen: int = field(metadata={"help": "Maximum sequence length."})128 129 max_batch_size: int = field(metadata={"help": "Maximum batch size."})130 131 seqlen_offset: int = field(default=0, metadata={"help": "Sequence length offset."})132 133 batch_size_offset: int = field(default=0, metadata={"help": "Batch size offset."})134 135 key_value_memory_dict: Dict[str, Any] = field(136 default_factory=dict, metadata={"help": "Key value memory dictionary."}137 )138 139 lengths_per_sample: torch.Tensor = field(140 default=None, metadata={"help": "Lengths per sample."}141 )142 143 144class Embedding(nn.Module):145 """Token embedding with dropout."""146 147 def __init__(self, config: PretrainedConfig) -> None:148 super().__init__()149 150 self.wte = nn.Embedding(config.vocab_size, config.n_embd)151 self.drop = nn.Dropout(config.embd_pdrop)152 153 def forward(self, input_ids: torch.LongTensor) -> torch.FloatTensor:154 input_shape = input_ids.size()155 input_ids = input_ids.view(-1, input_shape[-1])156 157 hidden_states = self.wte(input_ids)158 hidden_states = self.drop(hidden_states)159 160 return hidden_states161 162 163# @torch.compile164def _apply_rotary_emb(165 x: torch.FloatTensor,166 cos: torch.FloatTensor,167 sin: torch.FloatTensor,168) -> torch.FloatTensor:169 _, seqlen, _, _ = x.shape170 _, rotary_dim = cos.shape171 rotary_dim *= 2172 173 x_rot = x[:, :, :, :rotary_dim]174 x_pass = x[:, :, :, rotary_dim:]175 176 x1, x2 = x_rot.chunk(2, dim=-1)177 c, s = rearrange(cos[:seqlen], "s d -> s 1 d"), rearrange(178 sin[:seqlen], "s d -> s 1 d"179 )180 x1, x2, c, s = [t.to(dtype=torch.float32) for t in [x1, x2, c, s]]181 182 x_rot = torch.cat([x1 * c - x2 * s, x1 * s + x2 * c], axis=-1).to(x.dtype)183 184 return torch.cat([x_rot, x_pass], axis=-1)185 186 187# @torch.compile188def _apply_rotary_emb_kv(189 kv: torch.FloatTensor,190 cos: torch.FloatTensor,191 sin: torch.FloatTensor,192 cos_k: Optional[torch.FloatTensor] = None,193 sin_k: Optional[torch.FloatTensor] = None,194) -> torch.FloatTensor:195 _, seqlen, _, _, _ = kv.shape196 _, rotary_dim = cos.shape197 rotary_dim *= 2198 199 k_rot = kv[:, :, 0, :, :rotary_dim]200 k_pass = kv[:, :, 0, :, rotary_dim:]201 202 k1, k2 = k_rot.chunk(2, dim=-1)203 c, s = rearrange(cos[:seqlen], "s d -> s 1 d"), rearrange(204 sin[:seqlen], "s d -> s 1 d"205 )206 k1, k2, c, s = [t.to(dtype=torch.float32) for t in [k1, k2, c, s]]207 208 k_rot = torch.cat([k1 * c - k2 * s, k1 * s + k2 * c], axis=-1).to(kv.dtype)209 210 return torch.cat(211 [212 torch.cat([k_rot, k_pass], axis=-1).unsqueeze(2),213 kv[:, :, 1:2, :, :],214 ],215 axis=2,216 )217 218 219# @torch.compile220def _apply_rotary_emb_qkv(221 qkv: torch.FloatTensor,222 cos: torch.FloatTensor,223 sin: torch.FloatTensor,224 cos_k: Optional[torch.FloatTensor] = None,225 sin_k: Optional[torch.FloatTensor] = None,226) -> torch.FloatTensor:227 _, seqlen, _, _, _ = qkv.shape228 _, rotary_dim = cos.shape229 rotary_dim *= 2230 231 q_rot = qkv[:, :, 0, :, :rotary_dim]232 q_pass = qkv[:, :, 0, :, rotary_dim:]233 234 k_rot = qkv[:, :, 1, :, :rotary_dim]235 k_pass = qkv[:, :, 1, :, rotary_dim:]236 237 q1, q2 = q_rot.chunk(2, dim=-1)238 k1, k2 = k_rot.chunk(2, dim=-1)239 c, s = rearrange(cos[:seqlen], "s d -> s 1 d"), rearrange(240 sin[:seqlen], "s d -> s 1 d"241 )242 q1, q2, k1, k2, c, s = [t.to(dtype=torch.float32) for t in [q1, q2, k1, k2, c, s]]243 244 q_rot = torch.cat([q1 * c - q2 * s, q1 * s + q2 * c], axis=-1).to(qkv.dtype)245 k_rot = torch.cat([k1 * c - k2 * s, k1 * s + k2 * c], axis=-1).to(qkv.dtype)246 247 return torch.cat(248 [249 torch.cat([q_rot, q_pass], axis=-1).unsqueeze(2),250 torch.cat([k_rot, k_pass], axis=-1).unsqueeze(2),251 qkv[:, :, 2:3, :, :],252 ],253 axis=2,254 )255 256 257class RotaryEmbedding(nn.Module):258 """Rotary positional embedding (RoPE).259 260 Reference:261 RoFormer: Enhanced Transformer with Rotary Position Embedding.262 https://arxiv.org/pdf/2104.09864.pdf.263 264 """265 266 def __init__(267 self,268 dim: int,269 base: int = 10000,270 scale_base: Optional[float] = None,271 pos_idx_in_fp32: bool = True,272 max_position_embeddings: int = 2048,273 device: Optional[str] = None,274 **kwargs,275 ) -> None:276 super().__init__()277 278 if scale_base is not None:279 raise NotImplementedError280 281 self.dim = dim282 self.base = float(base)283 self.scale_base = scale_base284 self.pos_idx_in_fp32 = pos_idx_in_fp32285 self.max_position_embeddings = max_position_embeddings286 self.device = device287 288 # Generate and save the inverse frequency buffer (non-trainable)289 inv_freq = self._compute_inv_freq(device)290 self.register_buffer("inv_freq", inv_freq, persistent=False)291 292 # Generate and save the scale buffer (non-trainable)293 scale = (294 (torch.arange(0, dim, 2, device=device, dtype=torch.float32) + 0.4 * dim)295 / (1.4 * dim)296 if scale_base is not None297 else None298 )299 self.register_buffer("scale", scale, persistent=False)300 301 # Initialize cached attributes since ONNX can't rely on dynamic initialization302 self._update_cos_sin_cache(303 max_position_embeddings, device=device, dtype=torch.float32304 )305 306 def _compute_inv_freq(self, device: Optional[str] = None) -> torch.FloatTensor:307 return 1.0 / (308 self.base309 ** (310 torch.arange(0, self.dim, 2, device=device, dtype=torch.float32)311 / self.dim312 )313 )314 315 def _update_cos_sin_cache(316 self,317 seqlen: int,318 device: Optional[str] = None,319 dtype: Optional[torch.dtype] = None,320 ) -> None:321 self._seq_len_cached = seqlen322 323 # fp32 is preferred since the output of `torch.arange` can be quite large324 # and bf16 would lose a lot of precision325 if self.pos_idx_in_fp32:326 t = torch.arange(seqlen, device=device, dtype=torch.float32)327 if self.inv_freq.dtype != torch.float32:328 inv_freq = self._compute_inv_freq(device=device)329 else:330 inv_freq = self.inv_freq331 else:332 t = torch.arange(seqlen, device=device, dtype=self.inv_freq.dtype)333 inv_freq = self.inv_freq334 335 # `torch.outer` is preferred since `torch.einsum` converts from fp32 to fp16 if used with AMP336 freqs = torch.outer(t, inv_freq)337 if self.scale is None:338 self._cos_cached = torch.cos(freqs).to(dtype)339 self._sin_cached = torch.sin(freqs).to(dtype)340 else:341 power = (342 torch.arange(seqlen, dtype=self.scale.dtype, device=self.scale.device)343 - seqlen // 2344 ) / self.scale_base345 scale = self.scale.to(device=power.device) ** rearrange(power, "s -> s 1")346 347 # Force the scale multiplication to happen in fp32348 self._cos_cached = (torch.cos(freqs) * scale).to(dtype)349 self._sin_cached = (torch.sin(freqs) * scale).to(dtype)350 self._cos_k_cached = (torch.cos(freqs) / scale).to(dtype)351 self._sin_k_cached = (torch.sin(freqs) / scale).to(dtype)352 353 def forward(354 self,355 qkv: torch.Tensor,356 kv: Optional[torch.Tensor] = None,357 seqlen_offset: int = 0,358 **kwargs,359 ) -> Tuple[torch.Tensor, torch.Tensor]:360 if (361 self._seq_len_cached < qkv.shape[1] + seqlen_offset362 or self._cos_cached.device != qkv.device363 or self._cos_cached.dtype != qkv.dtype364 or (self.training and self._cos_cached.is_inference())365 ):366 self._update_cos_sin_cache(367 qkv.shape[1] + seqlen_offset, device=qkv.device, dtype=qkv.dtype368 )369 370 if kv is None:371 return _apply_rotary_emb_qkv(372 qkv,373 self._cos_cached[seqlen_offset:],374 self._sin_cached[seqlen_offset:],375 )376 else:377 q = _apply_rotary_emb(378 qkv,379 self._cos_cached[seqlen_offset:],380 self._sin_cached[seqlen_offset:],381 )382 kv = _apply_rotary_emb_kv(383 kv,384 self._cos_cached[seqlen_offset:],385 self._sin_cached[seqlen_offset:],386 )387 388 return q, kv389 390 391class MLP(nn.Module):392 """Multi-Layer Perceptron.393 394 Reference:395 Attention Is All You Need.396 https://arxiv.org/pdf/1706.03762.pdf.397 398 """399 400 def __init__(401 self,402 config: PretrainedConfig,403 n_inner: Optional[int] = None,404 act_fn: Optional[str] = None,405 ) -> None:406 super().__init__()407 408 act_fn = config.activation_function if act_fn is None else act_fn409 410 n_inner = getattr(config, "n_inner", None) if n_inner is None else n_inner411 n_inner = n_inner if n_inner is not None else 4 * config.n_embd412 413 self.fc1 = nn.Linear(config.n_embd, n_inner)414 self.fc2 = nn.Linear(n_inner, config.n_embd)415 self.act = ACT2FN[act_fn]416 417 def forward(self, hidden_states: torch.FloatTensor) -> torch.FloatTensor:418 hidden_states = self.fc1(hidden_states)419 hidden_states = self.act(hidden_states)420 hidden_states = self.fc2(hidden_states)421 422 return hidden_states423 424 425class SelfAttention(nn.Module):426 """Self-attention layer (compatible with PyTorch).427 428 Reference:429 https://github.com/Dao-AILab/flash-attention/blob/main/flash_attn/modules/mha.py.430 431 """432 433 def __init__(434 self,435 causal: bool = True,436 softmax_scale: Optional[float] = None,437 attention_dropout: float = 0.0,438 ) -> None:439 super().__init__()440 441 self.causal = causal442 self.softmax_scale = softmax_scale443 self.drop = nn.Dropout(attention_dropout)444 445 @torch.autocast("cpu", enabled=False)446 @torch.autocast("cuda", enabled=False)447 def forward(448 self,449 qkv: torch.FloatTensor,450 causal: bool = None,451 key_padding_mask: Optional[torch.BoolTensor] = None,452 **kwargs,453 ) -> torch.FloatTensor:454 batch_size, seqlen = qkv.shape[0], qkv.shape[1]455 q, k, v = qkv.unbind(dim=2)456 457 q = q.to(torch.float32)458 k = k.to(torch.float32)459 460 causal = self.causal if causal is None else causal461 softmax_scale = self.softmax_scale or 1.0 / math.sqrt(q.shape[-1])462 463 # Autocast is manually disabled to avoid `torch.einsum` performing the operation464 # using float16, which might lead to overflow465 scores = torch.einsum("bthd,bshd->bhts", q, k * softmax_scale)466 467 if key_padding_mask is not None:468 padding_mask = torch.full(469 (batch_size, seqlen), -10000.0, dtype=scores.dtype, device=scores.device470 )471 padding_mask.masked_fill_(key_padding_mask, 0.0)472 473 scores = scores + rearrange(padding_mask, "b s -> b 1 1 s")474 475 if causal:476 causal_mask = torch.triu(477 torch.full((seqlen, seqlen), -10000.0, device=scores.device), 1478 )479 scores = scores + causal_mask.to(dtype=scores.dtype)480 481 attention = torch.softmax(scores, dim=-1).to(v.dtype)482 attention = self.drop(attention)483 484 output = torch.einsum("bhts,bshd->bthd", attention, v)485 486 return output487 488 489class CrossAttention(nn.Module):490 """Cross-attention layer (compatible with PyTorch).491 492 Reference:493 https://github.com/Dao-AILab/flash-attention/blob/main/flash_attn/modules/mha.py.494 495 """496 497 def __init__(498 self,499 causal: bool = True,500 softmax_scale: Optional[float] = None,501 attention_dropout: float = 0.0,502 ) -> None:503 super().__init__()504 505 self.causal = causal506 self.softmax_scale = softmax_scale507 self.drop = nn.Dropout(attention_dropout)508 509 @torch.autocast("cpu", enabled=False)510 @torch.autocast("cuda", enabled=False)511 def forward(512 self,513 q: torch.FloatTensor,514 kv: torch.FloatTensor,515 causal: bool = None,516 key_padding_mask: Optional[torch.BoolTensor] = None,517 **kwargs,518 ) -> torch.FloatTensor:519 batch_size, seqlen_q = q.shape[0], q.shape[1]520 seqlen_k = kv.shape[1]521 522 if kv.shape[3] != q.shape[2]:523 kv = repeat(kv, "... hkv d -> ... (hkv g) d", g=q.shape[2] // kv.shape[3])524 k, v = kv.unbind(dim=2)525 526 q = q.to(torch.float32)527 k = k.to(torch.float32)528 529 causal = self.causal if causal is None else causal530 softmax_scale = self.softmax_scale or 1.0 / math.sqrt(q.shape[-1])531 532 # Autocast is manually disabled to avoid `torch.einsum` performing the operation533 # using float16, which might lead to overflow534 scores = torch.einsum("bthd,bshd->bhts", q, k * softmax_scale)535 536 if key_padding_mask is not None:537 padding_mask = torch.full(538 (batch_size, seqlen_k),539 -10000.0,540 dtype=scores.dtype,541 device=scores.device,542 )543 padding_mask.masked_fill_(key_padding_mask, 0.0)544 545 scores = scores + rearrange(padding_mask, "b s -> b 1 1 s")546 547 if causal:548 rows = rearrange(549 torch.arange(seqlen_q, device=q.device, dtype=torch.long), "s -> s 1"550 )551 cols = torch.arange(seqlen_k, device=k.device, dtype=torch.long)552 causal_mask = cols > rows + seqlen_k - seqlen_q553 554 scores = scores.masked_fill(causal_mask, -10000.0)555 556 attention = torch.softmax(scores, dim=-1).to(v.dtype)557 attention = self.drop(attention)558 559 output = torch.einsum("bhts,bshd->bthd", attention, v)560 561 return output562 563 564def _find_mha_dims(565 config: PretrainedConfig,566 n_head: Optional[int] = None,567 n_head_kv: Optional[int] = None,568 head_dim: Optional[int] = None,569) -> Tuple[int, int]:570 if n_head is None and head_dim is None:571 head_dim = config.n_embd // config.n_head572 n_head = config.n_head573 elif n_head is None or head_dim is None:574 raise ValueError("`n_head` and `head_dim` must be both specified or `None`.")575 576 if n_head_kv is None:577 n_head_kv = getattr(config, "n_head_kv", None) or n_head578 579 return n_head, n_head_kv, head_dim580 581 582def _update_kv_cache(583 kv: torch.FloatTensor, inference_params: InferenceParams, layer_idx: int584) -> torch.FloatTensor:585 num_heads, head_dim = kv.shape[-2:]586 587 if layer_idx not in inference_params.key_value_memory_dict:588 inference_params.key_value_memory_dict[layer_idx] = torch.empty(589 inference_params.max_batch_size,590 inference_params.max_seqlen,591 2,592 num_heads,593 head_dim,594 dtype=kv.dtype,595 device=kv.device,596 )597 598 batch_start = inference_params.batch_size_offset599 batch_end = batch_start + kv.shape[0]600 601 sequence_start = inference_params.seqlen_offset602 sequence_end = sequence_start + kv.shape[1]603 604 # When the current sequence length is equal to or larger than the maximum sequence length,605 # we need to concatenate the current `kv` with the cached `kv` to expand its length606 if sequence_end >= inference_params.max_seqlen:607 inference_params.key_value_memory_dict[layer_idx] = torch.concatenate(608 (inference_params.key_value_memory_dict[layer_idx], kv), dim=1609 )610 611 inference_params.key_value_memory_dict[layer_idx][612 batch_start:batch_end, sequence_start:sequence_end, ...613 ] = kv614 kv = inference_params.key_value_memory_dict[layer_idx][615 batch_start:batch_end, :sequence_end, ...616 ]617 618 return kv619 620 621class MHA(nn.Module):622 """Multi-head attention layer."""623 624 def __init__(625 self,626 config: PretrainedConfig,627 dtype: Optional[torch.dtype] = None,628 device: Optional[str] = None,629 rotary_dim: Optional[int] = None,630 rotary_base: float = 10000.0,631 rotary_scale_base: Optional[float] = None,632 n_head: Optional[int] = None,633 n_head_kv: Optional[int] = None,634 head_dim: Optional[int] = None,635 bias: bool = True,636 causal: bool = True,637 softmax_scale: Optional[float] = None,638 layer_idx: Optional[int] = None,639 return_residual: bool = False,640 checkpointing: bool = False,641 ) -> None:642 super().__init__()643 644 # Rotary embedding645 self.rotary_dim = (646 rotary_dim if rotary_dim is not None else getattr(config, "rotary_dim", 0)647 )648 649 if self.rotary_dim > 0:650 self.rotary_emb = RotaryEmbedding(651 self.rotary_dim,652 base=rotary_base,653 scale_base=rotary_scale_base,654 device=device,655 max_position_embeddings=config.n_positions,656 )657 658 # MLP659 self.n_head, self.n_head_kv, self.head_dim = _find_mha_dims(660 config, n_head=n_head, n_head_kv=n_head_kv, head_dim=head_dim661 )662 op_size = self.head_dim * (self.n_head + 2 * self.n_head_kv)663 hidden_size = config.n_embd664 665 linear_cls = FusedDense if config.fused_dense else nn.Linear666 if linear_cls is None:667 linear_cls = nn.Linear668 669 self.Wqkv = linear_cls(670 hidden_size, op_size, bias=bias, device=device, dtype=dtype671 )672 self.out_proj = linear_cls(673 hidden_size, hidden_size, bias=bias, device=device, dtype=dtype674 )675 676 # Attention677 self.inner_attn = SelfAttention(678 causal=causal,679 softmax_scale=softmax_scale,680 attention_dropout=config.attn_pdrop,681 )682 self.inner_cross_attn = CrossAttention(683 causal=causal,684 softmax_scale=softmax_scale,685 attention_dropout=config.attn_pdrop,686 )687 688 self.layer_idx = layer_idx689 self.return_residual = return_residual690 self.checkpointing = checkpointing691 692 def _forward_self_attn(693 self, x: torch.FloatTensor, key_padding_mask: Optional[torch.BoolTensor]694 ) -> torch.FloatTensor:695 qkv = self.Wqkv(x)696 qkv = rearrange(697 qkv, "... (three h d) -> ... three h d", three=3, d=self.head_dim698 )699 700 if self.rotary_dim > 0:701 qkv = self.rotary_emb(qkv)702 703 if self.checkpointing:704 return torch.utils.checkpoint.checkpoint(705 self.inner_attn, qkv, key_padding_mask=key_padding_mask706 )707 708 return self.inner_attn(qkv, key_padding_mask=key_padding_mask)709 710 def _forward_cross_attn(711 self,712 x: torch.FloatTensor,713 past_key_values: Optional[InferenceParams],714 key_padding_mask: Optional[torch.BoolTensor],715 ) -> torch.FloatTensor:716 batch_size = x.shape[0]717 718 qkv = self.Wqkv(x)719 720 q = qkv[..., : self.n_head * self.head_dim]721 q = rearrange(q, "... (h d) -> ... h d", d=self.head_dim)722 723 kv = qkv[..., self.n_head * self.head_dim :]724 kv = rearrange(kv, "... (two hkv d) -> ... two hkv d", two=2, d=self.head_dim)725 726 seqlen_offset = (727 past_key_values.seqlen_offset if past_key_values is not None else 0728 )729 causal = None if seqlen_offset == 0 else False730 if self.rotary_dim > 0:731 q, kv = self.rotary_emb(q, kv=kv, seqlen_offset=seqlen_offset)732 733 if past_key_values is not None:734 kv = _update_kv_cache(kv, past_key_values, self.layer_idx)735 736 if self.checkpointing:737 return torch.utils.checkpoint.checkpoint(738 self.inner_cross_attn,739 q,740 kv,741 key_padding_mask=key_padding_mask,742 causal=causal,743 )744 745 return self.inner_cross_attn(746 q, kv, key_padding_mask=key_padding_mask, causal=causal747 )748 749 def forward(750 self,751 x: torch.FloatTensor,752 past_key_values: Optional[InferenceParams] = None,753 attention_mask: Optional[Union[torch.LongTensor, torch.BoolTensor]] = None,754 **kwargs,755 ) -> Tuple[torch.FloatTensor, torch.FloatTensor]:756 if attention_mask is not None:757 attention_mask = attention_mask.bool()758 else:759 attention_mask = None760 761 # MHA762 if self.n_head == self.n_head_kv:763 if past_key_values is None:764 # If `past_key_values` are not supplied, we run self-attention765 attn_output = self._forward_self_attn(x, attention_mask)766 else:767 # If `past_key_values` are supplied, it means that we might have cached values and768 # could take advantage of cross-attention769 attn_output = self._forward_cross_attn(770 x, past_key_values, attention_mask771 )772 # MQA / GQA773 else:774 # Regardless of `past_key_values` being supplied or not, it always use cross-attention775 # because `q` and `kv` lengths might be different776 attn_output = self._forward_cross_attn(x, past_key_values, attention_mask)777 778 output = rearrange(attn_output, "... h d -> ... (h d)")779 output = self.out_proj(output)780 781 return output if not self.return_residual else (output, x)782 783 784class ParallelBlock(nn.Module):785 """Parallel block.786 787 This block applies parallel mixer and MLP layers to the input (used in GPT-J and CodeGen).788 789 """790 791 def __init__(792 self,793 config: PretrainedConfig,794 block_idx: Optional[int] = None,795 ) -> None:796 super().__init__()797 798 self.ln = nn.LayerNorm(config.n_embd, eps=config.layer_norm_epsilon)799 self.resid_dropout = nn.Dropout(config.resid_pdrop)800 self.block_idx = block_idx801 802 self.mixer = MHA(config, layer_idx=block_idx)803 self.mlp = MLP(config)804 805 def forward(806 self,807 hidden_states: torch.FloatTensor,808 past_key_values: Optional[Union[torch.FloatTensor, InferenceParams]] = None,809 attention_mask: Optional[torch.BoolTensor] = None,810 **kwargs,811 ) -> torch.FloatTensor:812 residual = hidden_states813 hidden_states = self.ln(hidden_states)814 815 attn_outputs = self.mixer(816 hidden_states,817 past_key_values=past_key_values,818 attention_mask=attention_mask,819 )820 if isinstance(attn_outputs, tuple):821 attn_outputs = attn_outputs[0]822 823 attn_outputs = self.resid_dropout(attn_outputs)824 feed_forward_hidden_states = self.resid_dropout(self.mlp(hidden_states))825 826 hidden_states = attn_outputs + feed_forward_hidden_states + residual827 828 return hidden_states829 830 831class CausalLMHead(nn.Module):832 """Causal Language Modeling head.833 834 Reference:835 Improving Language Understanding by Generative Pre-Training.836 https://cdn.openai.com/research-covers/language-unsupervised/language_understanding_paper.pdf.837 838 """839 840 def __init__(self, config: PretrainedConfig) -> None:841 super().__init__()842 843 self.ln = nn.LayerNorm(config.n_embd, eps=config.layer_norm_epsilon)844 self.linear = nn.Linear(config.n_embd, config.vocab_size)845 846 def forward(self, hidden_states: torch.FloatTensor) -> torch.FloatTensor:847 hidden_states = self.ln(hidden_states)848 logits = self.linear(hidden_states).to(torch.float32)849 850 return logits851 852 853class CausalLMLoss(nn.Module):854 """Causal Language Modeling loss.855 856 Reference:857 Improving Language Understanding by Generative Pre-Training.858 https://cdn.openai.com/research-covers/language-unsupervised/language_understanding_paper.pdf.859 860 """861 862 def __init__(self, shift_labels: bool = True) -> None:863 super().__init__()864 865 self.shift_labels = shift_labels866 self.loss_fct = nn.CrossEntropyLoss()867 868 def forward(869 self, logits: torch.FloatTensor, labels: torch.LongTensor870 ) -> torch.FloatTensor:871 if self.shift_labels:872 logits = logits[..., :-1, :].contiguous()873 labels = labels[..., 1:].contiguous()874 875 loss = self.loss_fct(logits.view(-1, logits.size(-1)), labels.view(-1))876 877 return loss878 879 880class PhiPreTrainedModel(PreTrainedModel):881 """Phi pre-trained model."""882 883 config_class = PhiConfig884 base_model_prefix = "transformer"885 supports_gradient_checkpointing = False886 _no_split_modules = ["ParallelBlock"]887 888 def __init__(self, *inputs, **kwargs) -> None:889 super().__init__(*inputs, **kwargs)890 891 def prepare_inputs_for_generation(892 self,893 input_ids: torch.LongTensor = None,894 inputs_embeds: torch.FloatTensor = None,895 past_key_values: Optional[Union[torch.FloatTensor, InferenceParams]] = None,896 attention_mask: Optional[Union[torch.LongTensor, torch.BoolTensor]] = None,897 **kwargs,898 ) -> Dict[str, Any]:899 if inputs_embeds is not None:900 max_batch_size = inputs_embeds.shape[0]901 seqlen_offset = inputs_embeds.shape[1] + input_ids.shape[1] - 2902 elif input_ids is not None:903 max_batch_size = input_ids.shape[0]904 seqlen_offset = input_ids.shape[1] - 1905 else:906 raise ValueError(907 "You have to specify either `input_ids` or `inputs_embeds`."908 )909 910 args = {}911 912 if past_key_values is None or not (913 isinstance(past_key_values, InferenceParams)914 ):915 past_key_values = InferenceParams(916 max_seqlen=self.config.n_positions,917 max_batch_size=max_batch_size,918 seqlen_offset=0,919 batch_size_offset=0,920 key_value_memory_dict={},921 lengths_per_sample=None,922 )923 if inputs_embeds is not None:924 args = {"inputs_embeds": inputs_embeds}925 elif input_ids is not None:926 args = {"input_ids": input_ids}927 else:928 raise ValueError(929 "You have to specify either `input_ids` or `inputs_embeds`."930 )931 else:932 # Assume that `past_key_values` has cached all tokens up to the last token in `input_ids`933 past_key_values.seqlen_offset = seqlen_offset934 input_ids = input_ids[:, -1].unsqueeze(-1)935 args = {"input_ids": input_ids}936 937 return {938 **args,939 "past_key_values": past_key_values,940 "attention_mask": attention_mask,941 }942 943 944class PhiModel(PhiPreTrainedModel):945 """Phi model."""946 947 _keys_to_ignore_on_load_missing = [""]948 _keys_to_ignore_on_load_unexpected = [r"h\.\d+\.mlp.(fc_in|fc_out)\.(weight|bias)"]949 950 def __init__(self, config: PhiConfig) -> None:951 super().__init__(config)952 953 self.embd = Embedding(config)954 self.h = nn.ModuleList(955 [ParallelBlock(config, block_idx=i) for i in range(config.n_layer)]956 )957 self.gradient_checkpointing = config.gradient_checkpointing958 self.post_init()959 960 def get_input_embeddings(self) -> nn.Embedding:961 return self.embd.wte962 963 def set_input_embeddings(self, new_embeddings: nn.Embedding) -> None:964 self.embd.wte = new_embeddings965 966 def forward(967 self,968 input_ids: torch.LongTensor = None,969 inputs_embeds: torch.FloatTensor = None,970 past_key_values: Optional[Union[torch.FloatTensor, InferenceParams]] = None,971 attention_mask: Optional[torch.BoolTensor] = None,972 ) -> torch.FloatTensor:973 if input_ids is not None and inputs_embeds is not None:974 raise ValueError(975 "You cannot specify both `input_ids` and `inputs_embeds` at the same time."976 )977 elif input_ids is None and inputs_embeds is None:978 raise ValueError(979 "You have to specify either `input_ids` or `inputs_embeds`."980 )981 elif input_ids is not None:982 hidden_states = self.embd(input_ids)983 else:984 hidden_states = inputs_embeds985 986 for layer in self.h:987 if self.gradient_checkpointing:988 hidden_states = torch.utils.checkpoint.checkpoint(989 layer.__call__,990 hidden_states,991 past_key_values,992 attention_mask,993 use_reentrant=True,994 )995 else:996 hidden_states = layer(997 hidden_states,998 past_key_values=past_key_values,999 attention_mask=attention_mask,1000 )1001 1002 return hidden_states1003 1004 1005class PhiForCausalLM(PhiPreTrainedModel):1006 """Phi for Causal Language Modeling."""1007 1008 _keys_to_ignore_on_load_missing = [""]1009 _keys_to_ignore_on_load_unexpected = [1010 r"transformer\.h\.\d+\.mlp.(fc_in|fc_out)\.(weight|bias)"1011 ]1012 1013 def __init__(self, config: PhiConfig) -> None:1014 super().__init__(config)1015 1016 self.transformer = PhiModel(config)1017 self.lm_head = CausalLMHead(config)1018 self.loss = CausalLMLoss()1019 1020 self.post_init()1021 1022 def get_output_embeddings(self) -> nn.Linear:1023 return self.lm_head.linear1024 1025 def set_output_embeddings(self, new_embeddings: nn.Linear) -> None:1026 self.lm_head.linear = new_embeddings1027 1028 def forward(1029 self,1030 input_ids: torch.LongTensor = None,1031 inputs_embeds: torch.FloatTensor = None,1032 past_key_values: Optional[Union[torch.FloatTensor, InferenceParams]] = None,1033 attention_mask: Optional[torch.BoolTensor] = None,1034 labels: Optional[torch.LongTensor] = None,1035 **kwargs,1036 ) -> CausalLMOutputWithPast:1037 hidden_states = self.transformer(1038 input_ids,1039 inputs_embeds,1040 past_key_values=past_key_values,1041 attention_mask=attention_mask,1042 )1043 lm_logits = self.lm_head(hidden_states)1044 1045 loss = None1046 if labels is not None:1047 loss = self.loss(lm_logits, labels)1048 1049 return CausalLMOutputWithPast(1050 loss=loss, logits=lm_logits, past_key_values=past_key_values1051 )1052 1053 1054class VisionEncoder(nn.Module):1055 def __init__(self, model_path: str = "model") -> None:1056 super().__init__()1057 self.model = torch.jit.load(f"{model_path}/vision.pt").to(DEVICE, dtype=DTYPE)1058 self.preprocess = Compose(1059 [1060 Resize(size=(384, 384), interpolation=InterpolationMode.BICUBIC),1061 ToImage(),1062 ToDtype(torch.float32, scale=True),1063 Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]),1064 ]1065 )1066 1067 def __call__(self, image: Image) -> torch.Tensor:1068 with torch.no_grad():1069 image_vec = self.preprocess(image.convert("RGB")).unsqueeze(0)1070 image_vec = image_vec[:, :, :-6, :-6]1071 image_vec = rearrange(1072 image_vec, "b c (h p1) (w p2) -> b (h w) (c p1 p2)", p1=14, p2=141073 )1074 1075 image_vec = image_vec.to(DEVICE, dtype=DTYPE)1076 return self.model(image_vec)1077 1078 1079class TextModel(nn.Module):1080 def __init__(self, model_path: str = "model") -> None:1081 super().__init__()1082 self.tokenizer = Tokenizer.from_pretrained(f"{model_path}/tokenizer")1083 phi_config = PhiConfig.from_pretrained(f"{model_path}/text_model_cfg.json")1084 1085 with init_empty_weights():1086 self.model = PhiForCausalLM(phi_config)1087 1088 self.model = load_checkpoint_and_dispatch(1089 self.model,1090 f"{model_path}/text_model.pt",1091 device_map={"": DEVICE},1092 dtype=DTYPE,1093 )1094 1095 self.text_emb = self.model.get_input_embeddings()1096 1097 def input_embeds(self, prompt, image_embeds):1098 embeds = []1099 1100 def _add_toks(toks):1101 embeds.append(self.text_emb(toks))1102 1103 def _tokenize(txt):1104 return self.tokenizer(1105 txt, return_tensors="pt", add_special_tokens=False1106 ).input_ids.to(self.model.device)1107 1108 # Add BOS token1109 _add_toks(1110 torch.tensor([[self.tokenizer.bos_token_id]], device=self.model.device)1111 )1112 1113 if "<image>" not in prompt:1114 embeds.append(self.text_emb(_tokenize(prompt)))1115 else:1116 assert prompt.count("<image>") == 11117 before, after = prompt.split("<image>")1118 embeds.append(self.text_emb(_tokenize(f"{before}<image>")))1119 embeds.append(image_embeds.to(self.model.device))1120 embeds.append(self.text_emb(_tokenize(f"</image>{after}")))1121 1122 return torch.cat(embeds, dim=1)1123 1124 def generate(1125 self, image_embeds, prompt, eos_text="Human:", max_new_tokens=128, **kwargs1126 ):1127 eos_tokens = self.tokenizer(eos_text, add_special_tokens=False)[0].ids1128 1129 generate_config = {1130 "eos_token_id": eos_tokens,1131 "bos_token_id": self.tokenizer.bos_token_id,1132 "pad_token_id": self.tokenizer.eos_token_id,1133 "max_new_tokens": max_new_tokens,1134 **kwargs,1135 }1136 1137 with torch.no_grad():1138 inputs_embeds = self.input_embeds(prompt, image_embeds)1139 output_ids = self.model.generate(1140 inputs_embeds=inputs_embeds, **generate_config1141 )1142 1143 return self.tokenizer.batch_decode(output_ids, skip_special_tokens=True)1144 1145 def answer_question(self, image_embeds, question, **kwargs):1146 prompt = f"<image>\n\nQuestion: {question}\n\nAnswer:"1147 answer = self.generate(1148 image_embeds,1149 prompt,1150 eos_text="<END>",1151 max_new_tokens=128,1152 **kwargs,1153 )[0]1154 1155 return re.sub("<$", "", re.sub("END$", "", answer)).strip()1156 1157 1158##### GRADIO INTERFACE #####1159 1160import gradio as gr1161from huggingface_hub import snapshot_download1162from threading import Thread1163from transformers import TextIteratorStreamer1164import hashlib1165import os1166 1167model_path = snapshot_download("vikhyatk/moondream1")1168 1169vision_encoder = VisionEncoder(model_path).to(DEVICE, dtype=DTYPE)1170text_model = TextModel(model_path).to(DEVICE, dtype=DTYPE)1171 1172 1173def cached_vision_encoder(image):1174 # Calculate checksum of the image1175 image_hash = hashlib.sha256(image.tobytes()).hexdigest()1176 1177 # Check if `image_encoder_cache/{image_hash}.pt` exists, if so load and return it.1178 # Otherwise, save the encoded image to `image_encoder_cache/{image_hash}.pt` and return it.1179 cache_path = f"image_encoder_cache/{image_hash}.pt"1180 if os.path.exists(cache_path):1181 return torch.load(cache_path).to(DEVICE, dtype=DTYPE)1182 else:1183 image_vec = vision_encoder(image).to("cpu", dtype=torch.float16)1184 os.makedirs("image_encoder_cache", exist_ok=True)1185 torch.save(image_vec, cache_path)1186 return image_vec.to(DEVICE, dtype=DTYPE)1187 1188 1189@spaces.GPU1190def answer_question(image, question):1191 yield "Encoding image..."1192 1193 streamer = TextIteratorStreamer(text_model.tokenizer, skip_special_tokens=True)1194 generation_kwargs = dict(1195 image_embeds=cached_vision_encoder(image), question=question, streamer=streamer1196 )1197 thread = Thread(target=text_model.answer_question, kwargs=generation_kwargs)1198 thread.start()1199 1200 buffer = ""