Squashed commit of the following:

commit 0a847cba9d6a906f16866cef55abc3d08591d384
Author: vik <vikhyatk@gmail.com>
Date:   Wed Jan 31 12:31:10 2024 -0800

    bugfix

commit de750a63ec3546c6f2dccde2d90902f314ac1cf8
Author: vik <vikhyatk@gmail.com>
Date:   Tue Jan 30 18:33:40 2024 -0800

    clean up inference interface a bit

commit 2ee0cdad0c6c7e981bfe21fb5160319fe328bbd1
Author: vik <vikhyatk@gmail.com>
Date:   Tue Jan 30 17:20:12 2024 -0800

    load tokenizer from HF

commit 4c981021959e440b3d3d33a6cb4c7324349587e4
Author: vik <vikhyatk@gmail.com>
Date:   Tue Jan 30 17:07:57 2024 -0800

    bugfix

commit d5fa2f95fe1799e5191defa339e06bdfcc079bb8
Author: vik <vikhyatk@gmail.com>
Date:   Tue Jan 30 15:45:51 2024 -0800

    stop using torch.jit.script for the vision encoder

    the interface is a little awkward right now, will be fixed shortly

commit cc252f3f5fc54a4d1ea4a41c1849bd458e5eb515
Merge: 0ce7485 d310e37
Author: vik <vikhyatk@gmail.com>
Date:   Tue Jan 30 13:33:47 2024 -0800

    Merge pull request #34 from eltociear/patch-1

    Update README.md

commit 0ce7485481f664ed90e90d0e257856e61cfcdc44
Merge: 3f4815b 2329525
Author: vik <vikhyatk@gmail.com>
Date:   Tue Jan 30 13:32:49 2024 -0800

    Merge branch 'mazzzystar-main'

commit 232952569a3b2c46d19d0a3bee25032982f8c68d
Author: vik <vikhyatk@gmail.com>
Date:   Tue Jan 30 13:32:36 2024 -0800

    add missing newline

commit d310e3739743fee903dd1afd52c572c2b9a23370
Author: Ikko Eltociear Ashimine <eltociear@gmail.com>
Date:   Wed Jan 31 00:13:48 2024 +0900

    Update README.md

    Huggingface -> Hugging Face

commit 53b57032078233856fbca8e19c4e4bd5b7c69586
Merge: aad2073 3f4815b
Author: Ke Fang <myfancoo@qq.com>
Date:   Tue Jan 30 17:19:20 2024 +0800

    Merge branch 'main' into main

commit aad20735513fceac5437a94d7308b55b95dcb216
Author: mazzzystar <2680461921@qq.com>
Date:   Tue Jan 30 17:16:29 2024 +0800

    change from Thread -> Queue to get the full response.

commit 6d497b787af56d56a3c307b1745936bcdc7ea9e9
Author: mazzzystar <2680461921@qq.com>
Date:   Tue Jan 30 17:11:18 2024 +0800

    resolve conflict for stream response.

commit e51173058648ad46875275aed8d6f8a7b595294d
Author: mazzzystar <2680461921@qq.com>
Date:   Tue Jan 30 17:02:53 2024 +0800

    resolve conflict for stream response.

commit 3f4815bd86aabb18724d74ef024adeff6c53914e
Author: vik <vikhyatk@gmail.com>
Date:   Tue Jan 30 00:50:30 2024 -0800

    bring back interactive mode text streaming

commit 948fa7ff49d990fa477b70c5768ad6a793edaf38
Author: mazzzystar <2680461921@qq.com>
Date:   Tue Jan 30 16:49:07 2024 +0800

    support for multi-round chat.

commit 4a2fb46a17ca333a5f53e3787e01bd727058395f
Author: vik <vikhyatk@gmail.com>
Date:   Mon Jan 29 21:50:18 2024 -0800

    simplify resize/cropping in the vision encoder

commit c206a5d4fd1afe0129f2ff6c0aa03bdbd372fddd
Author: vik <vikhyatk@gmail.com>
Date:   Mon Jan 29 21:45:03 2024 -0800

    detect and use appropriate device/dtype

commit 38af98596e59f2a6c25c6b52b2bd5a672dab4144
Merge: 15d5dd7 70c5358
Author: vik <vikhyatk@gmail.com>
Date:   Mon Jan 29 03:43:58 2024 -0800

    Merge pull request #31 from markusheimerl/main

    Refactored to reduce code size

commit 70c5358887477991d91241fd66b597a3afc44a33
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 09:40:45 2024 +0000

    Refactor code to improve readability and maintainability

commit 12f61ce90786c8d5664cbbd0d8b9701033b241ce
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 09:39:59 2024 +0000

    Refactor gradio_demo.py and vision_encoder.py

commit 11619b7ca7c920fda02625f53a78c0a13c31c945
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 09:34:17 2024 +0000

    Refactor PhiForCausalLM class in modeling_phi.py

commit 2bf4eb998166a326b6a244f7a92f1a853a0e7817
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 09:25:19 2024 +0000

    Refactor PhiForCausalLM class in modeling_phi.py

commit 78f2edd4dd7ede9004b17829b834dc4634de7f59
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 09:23:59 2024 +0000

    Remove unused imports and variables

commit e6252a05dfad4eb265089147cce6a8cc9404e0bc
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 09:22:49 2024 +0000

    Refactor PhiModel forward method

commit 83362add27945409abdd1e1e37dfe30ff674202c
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 09:22:10 2024 +0000

    Refactor return statement in modeling_phi.py

commit aea6dcc4e19a7ccc4d793e410377d712e11e46e6
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 09:18:51 2024 +0000

    Refactor PhiPreTrainedModel's input handling

commit f08caef1b82c8c6805d1500fdf55c4b7690113d8
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 09:16:06 2024 +0000

    Add MHA and ParallelBlock classes

commit 6be42cd2682038e1497ed7c8da1fa22049ec6b33
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 09:14:02 2024 +0000

    Simplified CausalLMHead implementation

commit 863f13265e12ff50f361aea874e26083bb8afc7f
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 09:13:00 2024 +0000

    Refactor ParallelBlock class in modeling_phi.py

commit 2d3d197546240e14a9454a1765d10a46b6ad1472
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 09:01:25 2024 +0000

    Refactor ParallelBlock class in modeling_phi.py

commit 77e19808f4e0b1f9d2e8fb072f7c9257ad351206
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 08:58:30 2024 +0000

    Refactor inner_cross_attn function call in MHA class

commit f7e4b670fbf2decc26a7f1d9241e1beeaa45f7ff
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 08:57:40 2024 +0000

    Refactor MHA forward method

commit d3ed301d3a292d6f9bed26af0babcb4a0c296abd
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 08:55:53 2024 +0000

    Update MHA class with rotary embeddings

commit 0afdd4054a1d691b41c4c0e52addf332af4810c3
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 08:54:17 2024 +0000

    Refactor self-attention forward pass in MHA class

commit 1a8dde4b28eeae5c9dd62a00f1402a8d8bf399ff
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 08:53:19 2024 +0000

    Refactor MHA initialization and simplify RotaryEmbedding usage

commit 7ce1f8c45a021361e87e36cb1c1c9e28e3714439
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 08:48:16 2024 +0000

    Refactor _update_kv_cache function to improve memory usage

commit 70bd853cb7df4b49adc161a60d1ff7749601dbd7
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 08:47:30 2024 +0000

    Refactor _find_mha_dims function signature

commit 3b9598ed5a6a1849c2b2b7b413ff03ec64809237
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 08:42:05 2024 +0000

    Refactor forward method in CrossAttention class

commit 0607a683e8bc00812f72c37858fb1622d7310556
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 08:37:54 2024 +0000

    Add Flash Attention module

commit 30fc465d1a056510167c3be05b3f7455e2f6e9e3
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 08:32:16 2024 +0000

    Add Flash Attention module

commit 66427ecb22d3f2d2cab0ce7dc3ffeeb919fd99ca
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 08:29:05 2024 +0000

    Refactor SelfAttention class in modeling_phi.py

commit 107f9e317fcbd8f3313fa5e00f01241f9e308339
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 08:26:29 2024 +0000

    Refactor MLP module in modeling_phi.py

commit bcdfaa6f9821fe5b8ddd28bad792029e8baab709
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 08:23:05 2024 +0000

    Refactor forward method in RotaryEmbedding class

commit 3f0b043d55a157217fd4a9dd86284c79f7602ef0
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 08:16:58 2024 +0000

    Refactor RotaryEmbedding class to improve precision and performance

commit fc71d3338c3188301f087a38fd9e14319cc74f15
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 08:13:58 2024 +0000

    Update import statements and remove unused code

commit 1be11477b5c9ad1fbd01f023737ad3801acf756e
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 08:12:09 2024 +0000

    Refactor RotaryEmbedding class to use non-trainable buffers

commit 89dde7df1a9237d0838d1577fa0cba3d5172bd20
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 08:05:15 2024 +0000

    Refactor rotary embedding functions in modeling_phi.py

commit 085ee6b292bc21221cda4b36615d5fcfebb3e7df
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 08:02:29 2024 +0000

    Refactor rotary embedding functions

commit 025c8d454e15355ef7e1137c69d1c28305a6be63
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 07:59:09 2024 +0000

    Refactor rotary embedding calculation in modeling_phi.py

commit 873b39219c1cbb2906e17fc4e92ca8da66bf800e
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 07:58:12 2024 +0000

    Refactor _apply_rotary_emb_kv function to improve readability and performance

commit 32382c2eaa988e714fcfc3b5526890a946d9a776
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 07:54:45 2024 +0000

    Refactor rotary embedding function in modeling_phi.py

commit 3226f9f55c0c93ba7f0138bcf4532f1a687cb3f7
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 07:44:18 2024 +0000

    Refactor InferenceParams and Embedding classes

commit d5b825d3c92bd8422dfe6d266f9a3c5e1c7f4c4c
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 07:24:26 2024 +0000

    Refactor PhiConfig class in configuration_phi.py

commit 7a4d6c34fa23fd7878c03349e38200bb6fed0d06
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 07:16:25 2024 +0000

    Refactor text_model.py and vision_encoder.py

commit ddf4a06c7171c9ee0a345fbd77b5bbe9c3cf6487
Author: Markus Heimerl <149831926+markusheimerl@users.noreply.github.com>
Date:   Mon Jan 29 07:11:37 2024 +0000

    Refactor code and remove unused imports
This commit is contained in:
Kijai
2024-02-01 12:47:00 +02:00
parent 9a93059224
commit 8cfbc3e4aa
10 changed files with 1079 additions and 1132 deletions
+1 -1
View File
@@ -6,7 +6,7 @@ a tiny vision language model that kicks ass and runs anywhere
1.6B parameter model built using SigLIP, Phi-1.5 and the LLaVA training dataset.
Weights are licensed under CC-BY-SA due to using the LLaVA dataset. Try it out
on [Huggingface Spaces](https://huggingface.co/spaces/vikhyatk/moondream1)!
on [Hugging Face Spaces](https://huggingface.co/spaces/vikhyatk/moondream1)!
**Benchmarks**
+2 -2
View File
@@ -1,2 +1,2 @@
from .vision_encoder import VisionEncoder
from .text_model import TextModel
from .util import detect_device
from .moondream import Moondream
@@ -1,31 +1,20 @@
# Copyright (c) Microsoft Corporation.
# Licensed under the MIT license.
import math
from typing import Optional
from transformers import PretrainedConfig
from typing import Optional
import math
class PhiConfig(PretrainedConfig):
"""Phi configuration."""
model_type = "phi-msft"
attribute_map = {
"max_position_embeddings": "n_positions",
"hidden_size": "n_embd",
"num_attention_heads": "n_head",
"num_hidden_layers": "n_layer",
}
def __init__(
self,
vocab_size: int = 50304,
vocab_size: int = 51200,
n_positions: int = 2048,
n_embd: int = 1024,
n_layer: int = 20,
n_embd: int = 2048,
n_layer: int = 24,
n_inner: Optional[int] = None,
n_head: int = 16,
n_head: int = 32,
n_head_kv: Optional[int] = None,
rotary_dim: Optional[int] = 32,
activation_function: Optional[str] = "gelu_new",
@@ -41,26 +30,45 @@ class PhiConfig(PretrainedConfig):
pad_vocab_size_multiple: int = 64,
gradient_checkpointing: bool = False,
**kwargs
) -> None:
self.vocab_size = int(
):
pad_vocab_size = (
math.ceil(vocab_size / pad_vocab_size_multiple) * pad_vocab_size_multiple
)
self.n_positions = n_positions
self.n_embd = n_embd
self.n_layer = n_layer
self.n_inner = n_inner
self.n_head = n_head
self.n_head_kv = n_head_kv
super().__init__(
vocab_size=pad_vocab_size,
n_positions=n_positions,
n_embd=n_embd,
n_layer=n_layer,
n_inner=n_inner,
n_head=n_head,
n_head_kv=n_head_kv,
activation_function=activation_function,
attn_pdrop=attn_pdrop,
embd_pdrop=embd_pdrop,
resid_pdrop=resid_pdrop,
layer_norm_epsilon=layer_norm_epsilon,
initializer_range=initializer_range,
pad_vocab_size_multiple=pad_vocab_size_multiple,
tie_word_embeddings=tie_word_embeddings,
gradient_checkpointing=gradient_checkpointing,
**kwargs
)
self.rotary_dim = min(rotary_dim, n_embd // n_head)
self.activation_function = activation_function
self.flash_attn = flash_attn
self.flash_rotary = flash_rotary
self.fused_dense = fused_dense
self.attn_pdrop = attn_pdrop
self.embd_pdrop = embd_pdrop
self.resid_pdrop = resid_pdrop
self.layer_norm_epsilon = layer_norm_epsilon
self.initializer_range = initializer_range
self.gradient_checkpointing = gradient_checkpointing
super().__init__(tie_word_embeddings=tie_word_embeddings, **kwargs)
attribute_map = {
"max_position_embeddings": "n_positions",
"hidden_size": "n_embd",
"num_attention_heads": "n_head",
"num_hidden_layers": "n_layer",
}
class MoondreamConfig(PretrainedConfig):
model_type = "moondream1"
def __init__(self, **kwargs):
self.phi_config = PhiConfig(**kwargs)
super().__init__(**kwargs)
+720
View File
@@ -0,0 +1,720 @@
# Copyright (c) Microsoft Corporation.
# Licensed under the MIT license.
#
# Copyright (c) 2022, Tri Dao, trid@cs.stanford.edu.
# Licensed under the BSD 3-Clause License.
from dataclasses import dataclass, field
from typing import Any, Dict, Optional, Union, Tuple
import math
import torch
import torch.nn as nn
from einops import rearrange, repeat
from transformers import PretrainedConfig, PreTrainedModel
from transformers.activations import ACT2FN
from transformers.modeling_outputs import CausalLMOutputWithPast
from .configuration_moondream import PhiConfig
FusedDense = None
@dataclass
class InferenceParams:
max_seqlen: int
max_batch_size: int
seqlen_offset: int = 0
batch_size_offset: int = 0
key_value_memory_dict: Dict[str, Any] = field(default_factory=dict)
lengths_per_sample: torch.Tensor = None
class Embedding(nn.Module):
def __init__(self, config: PretrainedConfig):
super().__init__()
self.wte = nn.Embedding(config.vocab_size, config.n_embd)
self.drop = nn.Dropout(config.embd_pdrop)
def forward(self, input_ids: torch.LongTensor) -> torch.FloatTensor:
return self.drop(self.wte(input_ids.view(-1, input_ids.size(-1))))
def _apply_rotary_emb(x, cos, sin):
seqlen, rotary_dim = x.size(1), cos.size(1) * 2
x_rot, x_pass = x[..., :rotary_dim], x[..., rotary_dim:]
x1, x2 = x_rot.chunk(2, dim=-1)
c, s = cos[:seqlen].unsqueeze(1), sin[:seqlen].unsqueeze(1)
x_rot = torch.cat([x1 * c - x2 * s, x1 * s + x2 * c], dim=-1)
return torch.cat([x_rot.to(x.dtype), x_pass], dim=-1)
def _apply_rotary_emb_kv(
kv: torch.FloatTensor, cos: torch.FloatTensor, sin: torch.FloatTensor
) -> torch.FloatTensor:
seqlen, rotary_dim = kv.shape[1], cos.shape[-1] * 2
k_rot = kv[:, :, 0, :, :rotary_dim].chunk(2, dim=-1)
k_pass = kv[:, :, 0, :, rotary_dim:]
c, s = cos[:seqlen].unsqueeze(1), sin[:seqlen].unsqueeze(1)
k_rot = torch.cat(
[k_rot[0] * c - k_rot[1] * s, k_rot[0] * s + k_rot[1] * c], dim=-1
)
return torch.cat(
[torch.cat([k_rot, k_pass], dim=-1).unsqueeze(2), kv[:, :, 1:2, :, :]], dim=2
)
def _apply_rotary_emb_qkv(
qkv: torch.FloatTensor, cos: torch.FloatTensor, sin: torch.FloatTensor
) -> torch.FloatTensor:
seqlen, rotary_dim = qkv.shape[1], cos.shape[1] * 2
c = cos[:seqlen].unsqueeze(1)
s = sin[:seqlen].unsqueeze(1)
qkv_rot = torch.stack(
[
torch.cat(
[
qkv[:, :, i, :, : rotary_dim // 2] * c
- qkv[:, :, i, :, rotary_dim // 2 : rotary_dim] * s,
qkv[:, :, i, :, : rotary_dim // 2] * s
+ qkv[:, :, i, :, rotary_dim // 2 : rotary_dim] * c,
],
dim=-1,
).to(qkv.dtype)
for i in range(2)
],
dim=2,
)
qkv_pass = qkv[:, :, :2, :, rotary_dim:].unsqueeze(2)
qkv_v = qkv[:, :, 2:3, :, :]
return torch.cat([qkv_rot, qkv_pass, qkv_v], dim=2)
class RotaryEmbedding(nn.Module):
# Enhanced Transformer with Rotary Position Embedding (https://arxiv.org/pdf/2104.09864.pdf)
def __init__(
self,
dim: int,
base: int = 10000,
scale_base: Optional[float] = None,
pos_idx_in_fp32: bool = True,
max_position_embeddings: int = 2048,
device: Optional[str] = None,
) -> None:
super().__init__()
# fp32 is preferred since the output of `torch.arange` can be quite large and bf16 would lose a lot of precision
self.dim, self.base, self.pos_idx_in_fp32, self.device = (
dim,
float(base),
pos_idx_in_fp32,
device,
)
self.max_position_embeddings = max_position_embeddings
if scale_base is not None:
raise NotImplementedError
# Generate and register the non-trainable buffers
self.register_buffer(
"inv_freq", self._compute_inv_freq(device), persistent=False
)
self.register_buffer(
"scale", self._calculate_scale(dim, scale_base, device), persistent=False
)
self._update_cos_sin_cache(
max_position_embeddings, device=device, dtype=torch.float32
)
def _calculate_scale(self, dim, scale_base, device):
return (
(
(
torch.arange(0, dim, 2, device=device, dtype=torch.float32)
+ 0.4 * dim
)
/ (1.4 * dim)
)
if scale_base is not None
else None
)
def _compute_inv_freq(self, device: Optional[str] = None) -> torch.FloatTensor:
return 1.0 / (
self.base
** (
torch.arange(0, self.dim, 2, device=device, dtype=torch.float32)
/ self.dim
)
)
def _update_cos_sin_cache(
self,
seqlen: int,
device: Optional[str] = None,
dtype: Optional[torch.dtype] = None,
) -> None:
self._seq_len_cached = seqlen
t = torch.arange(
seqlen,
device=device,
dtype=torch.float32 if self.pos_idx_in_fp32 else self.inv_freq.dtype,
)
inv_freq = (
self._compute_inv_freq(device=device)
if self.pos_idx_in_fp32 and self.inv_freq.dtype != torch.float32
else self.inv_freq
)
freqs = torch.outer(t, inv_freq)
def apply_scale(freqs, scale, operator, dtype):
result = operator(freqs)
return (result / scale).to(dtype) if scale is not None else result.to(dtype)
if scale := self.scale:
power = (
torch.arange(seqlen, dtype=scale.dtype, device=scale.device)
- seqlen // 2
) / self.scale_base
scale = scale.to(device=power.device) ** power.unsqueeze(1)
self._cos_cached = apply_scale(
freqs, 1 / scale if scale is not None else None, torch.cos, dtype
)
self._sin_cached = apply_scale(
freqs, 1 / scale if scale is not None else None, torch.sin, dtype
)
if scale is not None:
self._cos_k_cached = apply_scale(freqs, scale, torch.cos, dtype)
self._sin_k_cached = apply_scale(freqs, scale, torch.sin, dtype)
def forward(
self,
qkv: torch.Tensor,
kv: Optional[torch.Tensor] = None,
seqlen_offset: int = 0,
) -> Tuple[torch.Tensor, torch.Tensor]:
should_update = (
self._seq_len_cached < qkv.shape[1] + seqlen_offset
or self._cos_cached.device != qkv.device
or self._cos_cached.dtype != qkv.dtype
or (self.training and self._cos_cached.is_inference())
)
if should_update:
self._update_cos_sin_cache(
qkv.shape[1] + seqlen_offset, device=qkv.device, dtype=qkv.dtype
)
offset_cos = self._cos_cached[seqlen_offset:]
offset_sin = self._sin_cached[seqlen_offset:]
if kv is None:
return _apply_rotary_emb_qkv(qkv, offset_cos, offset_sin)
else:
return _apply_rotary_emb(qkv, offset_cos, offset_sin), _apply_rotary_emb_kv(
kv, offset_cos, offset_sin
)
class MLP(nn.Module):
def __init__(
self,
config: PretrainedConfig,
n_inner: Optional[int] = None,
act_fn: Optional[str] = None,
) -> None:
super().__init__()
n_inner = n_inner or getattr(config, "n_inner", None) or 4 * config.n_embd
act_fn = act_fn or config.activation_function
self.fc1 = nn.Linear(config.n_embd, n_inner)
self.fc2 = nn.Linear(n_inner, config.n_embd)
self.act = ACT2FN[act_fn]
def forward(self, hidden_states: torch.FloatTensor) -> torch.FloatTensor:
return self.fc2(self.act(self.fc1(hidden_states)))
# Flash Attention (https://github.com/Dao-AILab/flash-attention/blob/main/flash_attn/modules/mha.py)
class SelfAttention(nn.Module):
def __init__(
self,
causal: bool = True,
softmax_scale: Optional[float] = None,
attention_dropout: float = 0.0,
):
super().__init__()
self.causal = causal
self.softmax_scale = softmax_scale
self.drop = nn.Dropout(attention_dropout)
@torch.autocast("cpu", enabled=False)
@torch.autocast("cuda", enabled=False)
def forward(
self,
qkv: torch.FloatTensor,
causal: Optional[bool] = None,
key_padding_mask: Optional[torch.BoolTensor] = None,
):
q, k, v = qkv.chunk(3, dim=-1)
scale = self.softmax_scale or 1.0 / q.size(-1) ** 0.5
scores = (
torch.einsum("bthd,bshd->bhts", q.to(torch.float32), k.to(torch.float32))
* scale
)
if causal or self.causal:
scores.triu_(1).fill_(-10000.0)
if key_padding_mask is not None:
scores.masked_fill_(key_padding_mask[:, None, None, :], -10000.0)
attn = self.drop(torch.softmax(scores, dim=-1).to(v.dtype))
return torch.einsum("bhts,bshd->bthd", attn, v)
# Flash Attention (https://github.com/Dao-AILab/flash-attention/blob/main/flash_attn/modules/mha.py)
class CrossAttention(nn.Module):
def __init__(self, causal=True, softmax_scale=None, attention_dropout=0.0):
super().__init__()
self.causal = causal
self.softmax_scale = softmax_scale
self.drop = nn.Dropout(attention_dropout)
@torch.autocast("cpu", enabled=False)
@torch.autocast("cuda", enabled=False)
def forward(
self,
q: torch.FloatTensor,
kv: torch.FloatTensor,
causal: bool = None,
key_padding_mask: Optional[torch.BoolTensor] = None,
) -> torch.FloatTensor:
batch_size, seqlen_q = q.shape[0], q.shape[1]
seqlen_k = kv.shape[1]
if kv.shape[3] != q.shape[2]:
kv = repeat(kv, "... hkv d -> ... (hkv g) d", g=q.shape[2] // kv.shape[3])
k, v = kv.unbind(dim=2)
q = q.to(torch.float32)
k = k.to(torch.float32)
causal = self.causal if causal is None else causal
softmax_scale = self.softmax_scale or 1.0 / math.sqrt(q.shape[-1])
# Autocast is manually disabled to avoid `torch.einsum` performing the operation using float16, which might lead to overflow
scores = torch.einsum("bthd,bshd->bhts", q, k * softmax_scale)
if key_padding_mask is not None:
padding_mask = torch.full(
(batch_size, seqlen_k),
-10000.0,
dtype=scores.dtype,
device=scores.device,
)
padding_mask.masked_fill_(key_padding_mask, 0.0)
scores = scores + rearrange(padding_mask, "b s -> b 1 1 s")
if causal:
rows = rearrange(
torch.arange(seqlen_q, device=q.device, dtype=torch.long), "s -> s 1"
)
cols = torch.arange(seqlen_k, device=k.device, dtype=torch.long)
causal_mask = cols > rows + seqlen_k - seqlen_q
scores = scores.masked_fill(causal_mask, -10000.0)
attention = torch.softmax(scores, dim=-1).to(v.dtype)
attention = self.drop(attention)
output = torch.einsum("bhts,bshd->bthd", attention, v)
return output
def _find_mha_dims(
config: PretrainedConfig,
n_head: Optional[int] = None,
n_head_kv: Optional[int] = None,
head_dim: Optional[int] = None,
) -> Tuple[int, int]:
if n_head is None and head_dim is None:
head_dim = config.n_embd // config.n_head
n_head = config.n_head
elif n_head is None or head_dim is None:
raise ValueError("`n_head` and `head_dim` must be both specified or `None`.")
if n_head_kv is None:
n_head_kv = getattr(config, "n_head_kv", None) or n_head
return n_head, n_head_kv, head_dim
def _update_kv_cache(
kv: torch.FloatTensor, inference_params: InferenceParams, layer_idx: int
) -> torch.FloatTensor:
num_heads, head_dim = kv.shape[-2:]
layer_memory = inference_params.key_value_memory_dict.setdefault(
layer_idx,
torch.empty(
inference_params.max_batch_size,
inference_params.max_seqlen,
2,
num_heads,
head_dim,
dtype=kv.dtype,
device=kv.device,
),
)
batch_slice = slice(
inference_params.batch_size_offset,
inference_params.batch_size_offset + kv.shape[0],
)
seqlen_slice = slice(
inference_params.seqlen_offset, inference_params.seqlen_offset + kv.shape[1]
)
if seqlen_slice.stop >= inference_params.max_seqlen:
layer_memory = torch.cat((layer_memory, kv), dim=1)
inference_params.key_value_memory_dict[layer_idx] = layer_memory
layer_memory[batch_slice, seqlen_slice, ...] = kv
return layer_memory[batch_slice, : seqlen_slice.stop, ...]
# Multi-head attention layer with rotary embeddings
class MHA(nn.Module):
def __init__(
self,
config,
dtype=None,
device=None,
rotary_dim=None,
rotary_base=10000.0,
rotary_scale_base=None,
n_head=None,
n_head_kv=None,
head_dim=None,
bias=True,
causal=True,
softmax_scale=None,
layer_idx=None,
return_residual=False,
checkpointing=False,
):
super().__init__()
# Set rotary embedding if specified
self.rotary_dim = rotary_dim or getattr(config, "rotary_dim", 0)
if self.rotary_dim:
self.rotary_emb = RotaryEmbedding(
self.rotary_dim,
base=rotary_base,
scale_base=rotary_scale_base,
device=device,
max_position_embeddings=config.n_positions,
)
# Determine MHA dims from arguments or config
self.n_head, self.n_head_kv, self.head_dim = _find_mha_dims(
config, n_head, n_head_kv, head_dim
)
op_size = self.head_dim * (self.n_head + 2 * self.n_head_kv)
hidden_size = config.n_embd
# Choose Linear class based on config, FusedDense is optional
LinearClass = (
FusedDense if config.fused_dense and FusedDense is not None else nn.Linear
)
self.Wqkv = LinearClass(
hidden_size, op_size, bias=bias, device=device, dtype=dtype
)
self.out_proj = LinearClass(
hidden_size, hidden_size, bias=bias, device=device, dtype=dtype
)
# Initialize attention mechanisms
attn_kwargs = {
"causal": causal,
"softmax_scale": softmax_scale,
"attention_dropout": config.attn_pdrop,
}
self.inner_attn = SelfAttention(**attn_kwargs)
self.inner_cross_attn = CrossAttention(**attn_kwargs)
self.layer_idx = layer_idx
self.return_residual = return_residual
self.checkpointing = checkpointing
def _forward_self_attn(
self, x: torch.FloatTensor, key_padding_mask: Optional[torch.BoolTensor]
) -> torch.FloatTensor:
qkv = rearrange(
self.Wqkv(x), "... (three h d) -> ... three h d", three=3, d=self.head_dim
)
if self.rotary_dim > 0:
qkv = self.rotary_emb(qkv)
attn_func = (
torch.utils.checkpoint.checkpoint
if self.checkpointing
else lambda f, *args, **kwargs: f(*args, **kwargs)
)
return attn_func(self.inner_attn, qkv, key_padding_mask=key_padding_mask)
def _forward_cross_attn(
self,
x: torch.FloatTensor,
past_key_values: Optional[InferenceParams],
key_padding_mask: Optional[torch.BoolTensor],
) -> torch.FloatTensor:
qkv = self.Wqkv(x)
q, kv = (
qkv[..., : self.n_head * self.head_dim],
qkv[..., self.n_head * self.head_dim :],
)
q = rearrange(q, "... (h d) -> ... h d", d=self.head_dim)
kv = rearrange(kv, "... (two hkv d) -> ... two hkv d", two=2, d=self.head_dim)
seqlen_offset = (
past_key_values.seqlen_offset if past_key_values is not None else 0
)
causal = None if seqlen_offset == 0 else False
if self.rotary_dim > 0:
q, kv = self.rotary_emb(q, kv=kv, seqlen_offset=seqlen_offset)
if past_key_values is not None:
kv = _update_kv_cache(kv, past_key_values, self.layer_idx)
attn_func = (
torch.utils.checkpoint.checkpoint
if self.checkpointing
else lambda fn, *args, **kwargs: fn(*args, **kwargs)
)
return attn_func(
self.inner_cross_attn,
q,
kv,
key_padding_mask=key_padding_mask,
causal=causal,
)
def forward(
self,
x: torch.FloatTensor,
past_key_values: Optional[InferenceParams] = None,
attention_mask: Optional[Union[torch.LongTensor, torch.BoolTensor]] = None,
) -> Tuple[torch.FloatTensor, torch.FloatTensor]:
attention_mask = attention_mask.bool() if attention_mask is not None else None
use_cross_attn = self.n_head != self.n_head_kv or past_key_values is not None
attn_output_function = (
self._forward_cross_attn if use_cross_attn else self._forward_self_attn
)
attn_output = (
attn_output_function(x, past_key_values, attention_mask)
if use_cross_attn
else attn_output_function(x, attention_mask)
)
output = self.out_proj(rearrange(attn_output, "... h d -> ... (h d)"))
return (output, x) if self.return_residual else output
# Parallel block. This block applies parallel mixer and MLP layers to the input (used in GPT-J and CodeGen).
class ParallelBlock(nn.Module):
def __init__(self, config: PretrainedConfig, block_idx: Optional[int] = None):
super().__init__()
self.ln = nn.LayerNorm(config.n_embd, eps=config.layer_norm_epsilon)
self.resid_dropout = nn.Dropout(config.resid_pdrop)
self.block_idx = block_idx
self.mixer = MHA(config, layer_idx=block_idx)
self.mlp = MLP(config)
def forward(
self,
hidden_states: torch.FloatTensor,
past_key_values: Optional[Union[torch.FloatTensor, InferenceParams]] = None,
attention_mask: Optional[torch.BoolTensor] = None,
) -> torch.FloatTensor:
residual = hidden_states
hidden_states = self.ln(hidden_states)
attn_outputs = self.mixer(
hidden_states,
past_key_values=past_key_values,
attention_mask=attention_mask,
)
if isinstance(attn_outputs, tuple):
attn_outputs = attn_outputs[0]
attn_outputs = self.resid_dropout(attn_outputs)
feed_forward_hidden_states = self.resid_dropout(self.mlp(hidden_states))
return attn_outputs + feed_forward_hidden_states + residual
class CausalLMHead(nn.Module):
"""Causal Language Modeling head. Simplified version."""
def __init__(self, config):
super().__init__()
self.ln = nn.LayerNorm(config.n_embd, eps=config.layer_norm_epsilon)
self.linear = nn.Linear(config.n_embd, config.vocab_size)
def forward(self, hidden_states):
return self.linear(self.ln(hidden_states)).to(torch.float32)
# Improving Language Understanding by Generative Pre-Training
# (https://cdn.openai.com/research-covers/language-unsupervised/language_understanding_paper.pdf)
class CausalLMLoss(nn.Module):
def __init__(self, shift_labels: bool = True) -> None:
super().__init__()
self.shift_labels = shift_labels
self.loss_fct = nn.CrossEntropyLoss()
def forward(
self, logits: torch.FloatTensor, labels: torch.LongTensor
) -> torch.FloatTensor:
if self.shift_labels:
logits, labels = logits[..., :-1, :], labels[..., 1:]
return self.loss_fct(logits.reshape(-1, logits.size(-1)), labels.reshape(-1))
class PhiPreTrainedModel(PreTrainedModel):
config_class = PhiConfig
base_model_prefix = "transformer"
supports_gradient_checkpointing = False
_no_split_modules = ["ParallelBlock"]
def __init__(self, *inputs, **kwargs) -> None:
super().__init__(*inputs, **kwargs)
def prepare_inputs_for_generation(
self,
input_ids: torch.LongTensor = None,
inputs_embeds: torch.FloatTensor = None,
past_key_values: Optional[Union[torch.FloatTensor, InferenceParams]] = None,
attention_mask: Optional[Union[torch.LongTensor, torch.BoolTensor]] = None,
**kwargs,
) -> Dict[str, Any]:
if input_ids is None and inputs_embeds is None:
raise ValueError(
"You have to specify either `input_ids` or `inputs_embeds`."
)
max_batch_size = (
inputs_embeds.shape[0] if inputs_embeds is not None else input_ids.shape[0]
)
seqlen_offset = (
inputs_embeds.shape[1] + input_ids.shape[1] - 2
if inputs_embeds is not None
else input_ids.shape[1] - 1
)
args = (
{"inputs_embeds": inputs_embeds}
if inputs_embeds is not None
else {"input_ids": input_ids}
)
if not isinstance(past_key_values, InferenceParams):
past_key_values = InferenceParams(
max_seqlen=self.config.n_positions,
max_batch_size=max_batch_size,
seqlen_offset=0,
batch_size_offset=0,
key_value_memory_dict={},
lengths_per_sample=None,
)
else:
past_key_values.seqlen_offset = seqlen_offset
args = {"input_ids": input_ids[:, -1].unsqueeze(-1)}
return {
**args,
"past_key_values": past_key_values,
"attention_mask": attention_mask,
}
class PhiModel(PhiPreTrainedModel):
_keys_to_ignore_on_load_missing = [""]
_keys_to_ignore_on_load_unexpected = [r"h\.\d+\.mlp.(fc_in|fc_out)\.(weight|bias)"]
def __init__(self, config: PhiConfig) -> None:
super().__init__(config)
self.embd = Embedding(config)
self.h = nn.ModuleList(
[ParallelBlock(config, block_idx=i) for i in range(config.n_layer)]
)
self.gradient_checkpointing = config.gradient_checkpointing
self.post_init()
def get_input_embeddings(self) -> nn.Embedding:
return self.embd.wte
def set_input_embeddings(self, new_embeddings: nn.Embedding) -> None:
self.embd.wte = new_embeddings
def forward(
self,
input_ids: torch.LongTensor = None,
inputs_embeds: torch.FloatTensor = None,
past_key_values: Optional[Union[torch.FloatTensor, InferenceParams]] = None,
attention_mask: Optional[torch.BoolTensor] = None,
) -> torch.FloatTensor:
if (input_ids is None) == (inputs_embeds is None):
raise ValueError("Specify exactly one of `input_ids` or `inputs_embeds`.")
hidden_states = self.embd(input_ids) if input_ids is not None else inputs_embeds
for layer in self.h:
func = layer.__call__ if self.gradient_checkpointing else layer
args = (hidden_states, past_key_values, attention_mask)
hidden_states = (
torch.utils.checkpoint.checkpoint(func, *args, use_reentrant=True)
if self.gradient_checkpointing
else func(*args)
)
return hidden_states
class PhiForCausalLM(PhiPreTrainedModel):
_keys_to_ignore_on_load_missing, _keys_to_ignore_on_load_unexpected = (
[""],
[r"transformer\.h\.\d+\.mlp.(fc_in|fc_out)\.(weight|bias)"],
)
def __init__(self, config: PhiConfig) -> None:
super().__init__(config)
self.transformer = PhiModel(config)
self.lm_head = CausalLMHead(config)
self.loss = CausalLMLoss()
self.post_init()
def get_output_embeddings(self) -> nn.Linear:
return self.lm_head.linear
def set_output_embeddings(self, new_embeddings: nn.Linear) -> None:
self.lm_head.linear = new_embeddings
def forward(
self,
input_ids: torch.LongTensor = None,
inputs_embeds: torch.FloatTensor = None,
past_key_values: Optional[Union[torch.FloatTensor, InferenceParams]] = None,
attention_mask: Optional[torch.BoolTensor] = None,
labels: Optional[torch.LongTensor] = None,
**kwargs,
) -> CausalLMOutputWithPast:
hidden_states = self.transformer(
input_ids=input_ids,
inputs_embeds=inputs_embeds,
past_key_values=past_key_values,
attention_mask=attention_mask,
)
lm_logits = self.lm_head(hidden_states)
loss = self.loss(lm_logits, labels) if labels is not None else None
return CausalLMOutputWithPast(
loss=loss, logits=lm_logits, past_key_values=past_key_values
)
+100
View File
@@ -0,0 +1,100 @@
import torch
from .vision_encoder import VisionEncoder
from .text_model import TextModel
from .configuration_moondream import MoondreamConfig
from transformers import PreTrainedModel
import re
class Moondream(PreTrainedModel):
config_class = MoondreamConfig
def __init__(self, config):
super().__init__(config)
self.vision_encoder = VisionEncoder()
self.text_model = TextModel(config)
@property
def device(self):
return self.text_model.model.device
def encode_image(self, image):
return self.vision_encoder(image)
def input_embeds(self, prompt, image_embeds, tokenizer):
def _tokenize(txt):
return tokenizer(
txt, return_tensors="pt", add_special_tokens=False
).input_ids.to(self.device)
# Add BOS token
embeds = []
embeds.append(
self.text_model.text_emb(
(torch.tensor([[tokenizer.bos_token_id]], device=self.device))
)
)
if "<image>" not in prompt:
embeds.append(self.text_model.text_emb(_tokenize(prompt)))
else:
assert prompt.count("<image>") == 1
before, after = prompt.split("<image>")
embeds.append(self.text_model.text_emb(_tokenize(f"{before}<image>")))
embeds.append(image_embeds.to(self.device))
embeds.append(self.text_model.text_emb(_tokenize(f"</image>{after}")))
return torch.cat(embeds, dim=1)
def generate(
self,
image_embeds,
prompt,
tokenizer,
eos_text="Human:",
max_new_tokens=128,
**kwargs,
):
eos_tokens = tokenizer(eos_text, add_special_tokens=False)[0].ids
generate_config = {
"eos_token_id": eos_tokens,
"bos_token_id": tokenizer.bos_token_id,
"pad_token_id": tokenizer.eos_token_id,
"max_new_tokens": max_new_tokens,
**kwargs,
}
with torch.no_grad():
inputs_embeds = self.input_embeds(prompt, image_embeds, tokenizer)
output_ids = self.text_model.model.generate(
inputs_embeds=inputs_embeds, **generate_config
)
return tokenizer.batch_decode(output_ids, skip_special_tokens=True)
def answer_question(
self,
image_embeds,
question,
tokenizer,
chat_history="",
result_queue=None,
**kwargs,
):
prompt = f"<image>\n\n{chat_history}Question: {question}\n\nAnswer:"
answer = self.generate(
image_embeds,
prompt,
eos_text="<END>",
tokenizer=tokenizer,
max_new_tokens=128,
**kwargs,
)[0]
cleaned_answer = re.sub("<$", "", re.sub("END$", "", answer)).strip()
# Use the result_queue to pass the result if it is provided
if result_queue:
result_queue.put(cleaned_answer)
else:
return cleaned_answer
-970
View File
@@ -1,970 +0,0 @@
# Copyright (c) Microsoft Corporation.
# Licensed under the MIT license.
#
# Copyright (c) 2022, Tri Dao, trid@cs.stanford.edu.
# Licensed under the BSD 3-Clause License.
from __future__ import annotations
import math
from dataclasses import dataclass, field
from typing import Any, Dict, Optional, Tuple, Union
import torch
import torch.nn as nn
from einops import rearrange, repeat
from transformers import PretrainedConfig, PreTrainedModel
from transformers.activations import ACT2FN
from transformers.modeling_outputs import CausalLMOutputWithPast
from .configuration_phi import PhiConfig
pad_input, unpad_input = None, None
FlashRotaryEmbedding = None
FlashSelfAttention, FlashCrossAttention = None, None
FusedDense = None
@dataclass
class InferenceParams:
"""Inference parameters passed to model to efficiently calculate
and store context during inference.
Reference:
https://github.com/Dao-AILab/flash-attention/blob/main/flash_attn/utils/generation.py.
Args:
max_seqlen: Maximum sequence length.
max_batch_size: Maximum batch size.
seqlen_offset: Sequence length offset.
batch_size_offset: Batch size offset.
key_value_memory_dict: Key value memory dictionary.
lengths_per_sample: Lengths per sample.
"""
max_seqlen: int = field(metadata={"help": "Maximum sequence length."})
max_batch_size: int = field(metadata={"help": "Maximum batch size."})
seqlen_offset: int = field(default=0, metadata={"help": "Sequence length offset."})
batch_size_offset: int = field(default=0, metadata={"help": "Batch size offset."})
key_value_memory_dict: Dict[str, Any] = field(
default_factory=dict, metadata={"help": "Key value memory dictionary."}
)
lengths_per_sample: torch.Tensor = field(
default=None, metadata={"help": "Lengths per sample."}
)
class Embedding(nn.Module):
"""Token embedding with dropout."""
def __init__(self, config: PretrainedConfig) -> None:
super().__init__()
self.wte = nn.Embedding(config.vocab_size, config.n_embd)
self.drop = nn.Dropout(config.embd_pdrop)
def forward(self, input_ids: torch.LongTensor) -> torch.FloatTensor:
input_shape = input_ids.size()
input_ids = input_ids.view(-1, input_shape[-1])
hidden_states = self.wte(input_ids)
hidden_states = self.drop(hidden_states)
return hidden_states
# @torch.compile
def _apply_rotary_emb(
x: torch.FloatTensor,
cos: torch.FloatTensor,
sin: torch.FloatTensor,
) -> torch.FloatTensor:
_, seqlen, _, _ = x.shape
_, rotary_dim = cos.shape
rotary_dim *= 2
x_rot = x[:, :, :, :rotary_dim]
x_pass = x[:, :, :, rotary_dim:]
x1, x2 = x_rot.chunk(2, dim=-1)
c, s = rearrange(cos[:seqlen], "s d -> s 1 d"), rearrange(
sin[:seqlen], "s d -> s 1 d"
)
x1, x2, c, s = [t.to(dtype=torch.float32) for t in [x1, x2, c, s]]
x_rot = torch.cat([x1 * c - x2 * s, x1 * s + x2 * c], axis=-1).to(x.dtype)
return torch.cat([x_rot, x_pass], axis=-1)
# @torch.compile
def _apply_rotary_emb_kv(
kv: torch.FloatTensor,
cos: torch.FloatTensor,
sin: torch.FloatTensor,
cos_k: Optional[torch.FloatTensor] = None,
sin_k: Optional[torch.FloatTensor] = None,
) -> torch.FloatTensor:
_, seqlen, _, _, _ = kv.shape
_, rotary_dim = cos.shape
rotary_dim *= 2
k_rot = kv[:, :, 0, :, :rotary_dim]
k_pass = kv[:, :, 0, :, rotary_dim:]
k1, k2 = k_rot.chunk(2, dim=-1)
c, s = rearrange(cos[:seqlen], "s d -> s 1 d"), rearrange(
sin[:seqlen], "s d -> s 1 d"
)
k1, k2, c, s = [t.to(dtype=torch.float32) for t in [k1, k2, c, s]]
k_rot = torch.cat([k1 * c - k2 * s, k1 * s + k2 * c], axis=-1).to(kv.dtype)
return torch.cat(
[
torch.cat([k_rot, k_pass], axis=-1).unsqueeze(2),
kv[:, :, 1:2, :, :],
],
axis=2,
)
# @torch.compile
def _apply_rotary_emb_qkv(
qkv: torch.FloatTensor,
cos: torch.FloatTensor,
sin: torch.FloatTensor,
cos_k: Optional[torch.FloatTensor] = None,
sin_k: Optional[torch.FloatTensor] = None,
) -> torch.FloatTensor:
_, seqlen, _, _, _ = qkv.shape
_, rotary_dim = cos.shape
rotary_dim *= 2
q_rot = qkv[:, :, 0, :, :rotary_dim]
q_pass = qkv[:, :, 0, :, rotary_dim:]
k_rot = qkv[:, :, 1, :, :rotary_dim]
k_pass = qkv[:, :, 1, :, rotary_dim:]
q1, q2 = q_rot.chunk(2, dim=-1)
k1, k2 = k_rot.chunk(2, dim=-1)
c, s = rearrange(cos[:seqlen], "s d -> s 1 d"), rearrange(
sin[:seqlen], "s d -> s 1 d"
)
q1, q2, k1, k2, c, s = [t.to(dtype=torch.float32) for t in [q1, q2, k1, k2, c, s]]
q_rot = torch.cat([q1 * c - q2 * s, q1 * s + q2 * c], axis=-1).to(qkv.dtype)
k_rot = torch.cat([k1 * c - k2 * s, k1 * s + k2 * c], axis=-1).to(qkv.dtype)
return torch.cat(
[
torch.cat([q_rot, q_pass], axis=-1).unsqueeze(2),
torch.cat([k_rot, k_pass], axis=-1).unsqueeze(2),
qkv[:, :, 2:3, :, :],
],
axis=2,
)
class RotaryEmbedding(nn.Module):
"""Rotary positional embedding (RoPE).
Reference:
RoFormer: Enhanced Transformer with Rotary Position Embedding.
https://arxiv.org/pdf/2104.09864.pdf.
"""
def __init__(
self,
dim: int,
base: int = 10000,
scale_base: Optional[float] = None,
pos_idx_in_fp32: bool = True,
max_position_embeddings: int = 2048,
device: Optional[str] = None,
**kwargs,
) -> None:
super().__init__()
if scale_base is not None:
raise NotImplementedError
self.dim = dim
self.base = float(base)
self.scale_base = scale_base
self.pos_idx_in_fp32 = pos_idx_in_fp32
self.max_position_embeddings = max_position_embeddings
self.device = device
# Generate and save the inverse frequency buffer (non-trainable)
inv_freq = self._compute_inv_freq(device)
self.register_buffer("inv_freq", inv_freq, persistent=False)
# Generate and save the scale buffer (non-trainable)
scale = (
(torch.arange(0, dim, 2, device=device, dtype=torch.float32) + 0.4 * dim)
/ (1.4 * dim)
if scale_base is not None
else None
)
self.register_buffer("scale", scale, persistent=False)
# Initialize cached attributes since ONNX can't rely on dynamic initialization
self._update_cos_sin_cache(
max_position_embeddings, device=device, dtype=torch.float32
)
def _compute_inv_freq(self, device: Optional[str] = None) -> torch.FloatTensor:
return 1.0 / (
self.base
** (
torch.arange(0, self.dim, 2, device=device, dtype=torch.float32)
/ self.dim
)
)
def _update_cos_sin_cache(
self,
seqlen: int,
device: Optional[str] = None,
dtype: Optional[torch.dtype] = None,
) -> None:
self._seq_len_cached = seqlen
# fp32 is preferred since the output of `torch.arange` can be quite large
# and bf16 would lose a lot of precision
if self.pos_idx_in_fp32:
t = torch.arange(seqlen, device=device, dtype=torch.float32)
if self.inv_freq.dtype != torch.float32:
inv_freq = self._compute_inv_freq(device=device)
else:
inv_freq = self.inv_freq
else:
t = torch.arange(seqlen, device=device, dtype=self.inv_freq.dtype)
inv_freq = self.inv_freq
# `torch.outer` is preferred since `torch.einsum` converts from fp32 to fp16 if used with AMP
freqs = torch.outer(t, inv_freq)
if self.scale is None:
self._cos_cached = torch.cos(freqs).to(dtype)
self._sin_cached = torch.sin(freqs).to(dtype)
else:
power = (
torch.arange(seqlen, dtype=self.scale.dtype, device=self.scale.device)
- seqlen // 2
) / self.scale_base
scale = self.scale.to(device=power.device) ** rearrange(power, "s -> s 1")
# Force the scale multiplication to happen in fp32
self._cos_cached = (torch.cos(freqs) * scale).to(dtype)
self._sin_cached = (torch.sin(freqs) * scale).to(dtype)
self._cos_k_cached = (torch.cos(freqs) / scale).to(dtype)
self._sin_k_cached = (torch.sin(freqs) / scale).to(dtype)
def forward(
self,
qkv: torch.Tensor,
kv: Optional[torch.Tensor] = None,
seqlen_offset: int = 0,
**kwargs,
) -> Tuple[torch.Tensor, torch.Tensor]:
if (
self._seq_len_cached < qkv.shape[1] + seqlen_offset
or self._cos_cached.device != qkv.device
or self._cos_cached.dtype != qkv.dtype
or (self.training and self._cos_cached.is_inference())
):
self._update_cos_sin_cache(
qkv.shape[1] + seqlen_offset, device=qkv.device, dtype=qkv.dtype
)
if kv is None:
return _apply_rotary_emb_qkv(
qkv,
self._cos_cached[seqlen_offset:],
self._sin_cached[seqlen_offset:],
)
else:
q = _apply_rotary_emb(
qkv,
self._cos_cached[seqlen_offset:],
self._sin_cached[seqlen_offset:],
)
kv = _apply_rotary_emb_kv(
kv,
self._cos_cached[seqlen_offset:],
self._sin_cached[seqlen_offset:],
)
return q, kv
class MLP(nn.Module):
"""Multi-Layer Perceptron.
Reference:
Attention Is All You Need.
https://arxiv.org/pdf/1706.03762.pdf.
"""
def __init__(
self,
config: PretrainedConfig,
n_inner: Optional[int] = None,
act_fn: Optional[str] = None,
) -> None:
super().__init__()
act_fn = config.activation_function if act_fn is None else act_fn
n_inner = getattr(config, "n_inner", None) if n_inner is None else n_inner
n_inner = n_inner if n_inner is not None else 4 * config.n_embd
self.fc1 = nn.Linear(config.n_embd, n_inner)
self.fc2 = nn.Linear(n_inner, config.n_embd)
self.act = ACT2FN[act_fn]
def forward(self, hidden_states: torch.FloatTensor) -> torch.FloatTensor:
hidden_states = self.fc1(hidden_states)
hidden_states = self.act(hidden_states)
hidden_states = self.fc2(hidden_states)
return hidden_states
class SelfAttention(nn.Module):
"""Self-attention layer (compatible with PyTorch).
Reference:
https://github.com/Dao-AILab/flash-attention/blob/main/flash_attn/modules/mha.py.
"""
def __init__(
self,
causal: bool = True,
softmax_scale: Optional[float] = None,
attention_dropout: float = 0.0,
) -> None:
super().__init__()
self.causal = causal
self.softmax_scale = softmax_scale
self.drop = nn.Dropout(attention_dropout)
@torch.autocast("cpu", enabled=False)
@torch.autocast("cuda", enabled=False)
def forward(
self,
qkv: torch.FloatTensor,
causal: bool = None,
key_padding_mask: Optional[torch.BoolTensor] = None,
**kwargs,
) -> torch.FloatTensor:
batch_size, seqlen = qkv.shape[0], qkv.shape[1]
q, k, v = qkv.unbind(dim=2)
q = q.to(torch.float32)
k = k.to(torch.float32)
causal = self.causal if causal is None else causal
softmax_scale = self.softmax_scale or 1.0 / math.sqrt(q.shape[-1])
# Autocast is manually disabled to avoid `torch.einsum` performing the operation
# using float16, which might lead to overflow
scores = torch.einsum("bthd,bshd->bhts", q, k * softmax_scale)
if key_padding_mask is not None:
padding_mask = torch.full(
(batch_size, seqlen), -10000.0, dtype=scores.dtype, device=scores.device
)
padding_mask.masked_fill_(key_padding_mask, 0.0)
scores = scores + rearrange(padding_mask, "b s -> b 1 1 s")
if causal:
causal_mask = torch.triu(
torch.full((seqlen, seqlen), -10000.0, device=scores.device), 1
)
scores = scores + causal_mask.to(dtype=scores.dtype)
attention = torch.softmax(scores, dim=-1).to(v.dtype)
attention = self.drop(attention)
output = torch.einsum("bhts,bshd->bthd", attention, v)
return output
class CrossAttention(nn.Module):
"""Cross-attention layer (compatible with PyTorch).
Reference:
https://github.com/Dao-AILab/flash-attention/blob/main/flash_attn/modules/mha.py.
"""
def __init__(
self,
causal: bool = True,
softmax_scale: Optional[float] = None,
attention_dropout: float = 0.0,
) -> None:
super().__init__()
self.causal = causal
self.softmax_scale = softmax_scale
self.drop = nn.Dropout(attention_dropout)
@torch.autocast("cpu", enabled=False)
@torch.autocast("cuda", enabled=False)
def forward(
self,
q: torch.FloatTensor,
kv: torch.FloatTensor,
causal: bool = None,
key_padding_mask: Optional[torch.BoolTensor] = None,
**kwargs,
) -> torch.FloatTensor:
batch_size, seqlen_q = q.shape[0], q.shape[1]
seqlen_k = kv.shape[1]
if kv.shape[3] != q.shape[2]:
kv = repeat(kv, "... hkv d -> ... (hkv g) d", g=q.shape[2] // kv.shape[3])
k, v = kv.unbind(dim=2)
q = q.to(torch.float32)
k = k.to(torch.float32)
causal = self.causal if causal is None else causal
softmax_scale = self.softmax_scale or 1.0 / math.sqrt(q.shape[-1])
# Autocast is manually disabled to avoid `torch.einsum` performing the operation
# using float16, which might lead to overflow
scores = torch.einsum("bthd,bshd->bhts", q, k * softmax_scale)
if key_padding_mask is not None:
padding_mask = torch.full(
(batch_size, seqlen_k),
-10000.0,
dtype=scores.dtype,
device=scores.device,
)
padding_mask.masked_fill_(key_padding_mask, 0.0)
scores = scores + rearrange(padding_mask, "b s -> b 1 1 s")
if causal:
rows = rearrange(
torch.arange(seqlen_q, device=q.device, dtype=torch.long), "s -> s 1"
)
cols = torch.arange(seqlen_k, device=k.device, dtype=torch.long)
causal_mask = cols > rows + seqlen_k - seqlen_q
scores = scores.masked_fill(causal_mask, -10000.0)
attention = torch.softmax(scores, dim=-1).to(v.dtype)
attention = self.drop(attention)
output = torch.einsum("bhts,bshd->bthd", attention, v)
return output
def _find_mha_dims(
config: PretrainedConfig,
n_head: Optional[int] = None,
n_head_kv: Optional[int] = None,
head_dim: Optional[int] = None,
) -> Tuple[int, int]:
if n_head is None and head_dim is None:
head_dim = config.n_embd // config.n_head
n_head = config.n_head
elif n_head is None or head_dim is None:
raise ValueError("`n_head` and `head_dim` must be both specified or `None`.")
if n_head_kv is None:
n_head_kv = getattr(config, "n_head_kv", None) or n_head
return n_head, n_head_kv, head_dim
def _update_kv_cache(
kv: torch.FloatTensor, inference_params: InferenceParams, layer_idx: int
) -> torch.FloatTensor:
num_heads, head_dim = kv.shape[-2:]
if layer_idx not in inference_params.key_value_memory_dict:
inference_params.key_value_memory_dict[layer_idx] = torch.empty(
inference_params.max_batch_size,
inference_params.max_seqlen,
2,
num_heads,
head_dim,
dtype=kv.dtype,
device=kv.device,
)
batch_start = inference_params.batch_size_offset
batch_end = batch_start + kv.shape[0]
sequence_start = inference_params.seqlen_offset
sequence_end = sequence_start + kv.shape[1]
# When the current sequence length is equal to or larger than the maximum sequence length,
# we need to concatenate the current `kv` with the cached `kv` to expand its length
if sequence_end >= inference_params.max_seqlen:
inference_params.key_value_memory_dict[layer_idx] = torch.concatenate(
(inference_params.key_value_memory_dict[layer_idx], kv), dim=1
)
inference_params.key_value_memory_dict[layer_idx][
batch_start:batch_end, sequence_start:sequence_end, ...
] = kv
kv = inference_params.key_value_memory_dict[layer_idx][
batch_start:batch_end, :sequence_end, ...
]
return kv
class MHA(nn.Module):
"""Multi-head attention layer."""
def __init__(
self,
config: PretrainedConfig,
dtype: Optional[torch.dtype] = None,
device: Optional[str] = None,
rotary_dim: Optional[int] = None,
rotary_base: float = 10000.0,
rotary_scale_base: Optional[float] = None,
n_head: Optional[int] = None,
n_head_kv: Optional[int] = None,
head_dim: Optional[int] = None,
bias: bool = True,
causal: bool = True,
softmax_scale: Optional[float] = None,
layer_idx: Optional[int] = None,
return_residual: bool = False,
checkpointing: bool = False,
) -> None:
super().__init__()
# Rotary embedding
self.rotary_dim = (
rotary_dim if rotary_dim is not None else getattr(config, "rotary_dim", 0)
)
if self.rotary_dim > 0:
self.rotary_emb = RotaryEmbedding(
self.rotary_dim,
base=rotary_base,
scale_base=rotary_scale_base,
device=device,
max_position_embeddings=config.n_positions,
)
# MLP
self.n_head, self.n_head_kv, self.head_dim = _find_mha_dims(
config, n_head=n_head, n_head_kv=n_head_kv, head_dim=head_dim
)
op_size = self.head_dim * (self.n_head + 2 * self.n_head_kv)
hidden_size = config.n_embd
linear_cls = FusedDense if config.fused_dense else nn.Linear
if linear_cls is None:
linear_cls = nn.Linear
self.Wqkv = linear_cls(
hidden_size, op_size, bias=bias, device=device, dtype=dtype
)
self.out_proj = linear_cls(
hidden_size, hidden_size, bias=bias, device=device, dtype=dtype
)
# Attention
self.inner_attn = SelfAttention(
causal=causal,
softmax_scale=softmax_scale,
attention_dropout=config.attn_pdrop,
)
self.inner_cross_attn = CrossAttention(
causal=causal,
softmax_scale=softmax_scale,
attention_dropout=config.attn_pdrop,
)
self.layer_idx = layer_idx
self.return_residual = return_residual
self.checkpointing = checkpointing
def _forward_self_attn(
self, x: torch.FloatTensor, key_padding_mask: Optional[torch.BoolTensor]
) -> torch.FloatTensor:
qkv = self.Wqkv(x)
qkv = rearrange(
qkv, "... (three h d) -> ... three h d", three=3, d=self.head_dim
)
if self.rotary_dim > 0:
qkv = self.rotary_emb(qkv)
if self.checkpointing:
return torch.utils.checkpoint.checkpoint(
self.inner_attn, qkv, key_padding_mask=key_padding_mask
)
return self.inner_attn(qkv, key_padding_mask=key_padding_mask)
def _forward_cross_attn(
self,
x: torch.FloatTensor,
past_key_values: Optional[InferenceParams],
key_padding_mask: Optional[torch.BoolTensor],
) -> torch.FloatTensor:
batch_size = x.shape[0]
qkv = self.Wqkv(x)
q = qkv[..., : self.n_head * self.head_dim]
q = rearrange(q, "... (h d) -> ... h d", d=self.head_dim)
kv = qkv[..., self.n_head * self.head_dim :]
kv = rearrange(kv, "... (two hkv d) -> ... two hkv d", two=2, d=self.head_dim)
seqlen_offset = (
past_key_values.seqlen_offset if past_key_values is not None else 0
)
causal = None if seqlen_offset == 0 else False
if self.rotary_dim > 0:
q, kv = self.rotary_emb(q, kv=kv, seqlen_offset=seqlen_offset)
if past_key_values is not None:
kv = _update_kv_cache(kv, past_key_values, self.layer_idx)
if self.checkpointing:
return torch.utils.checkpoint.checkpoint(
self.inner_cross_attn,
q,
kv,
key_padding_mask=key_padding_mask,
causal=causal,
)
return self.inner_cross_attn(
q, kv, key_padding_mask=key_padding_mask, causal=causal
)
def forward(
self,
x: torch.FloatTensor,
past_key_values: Optional[InferenceParams] = None,
attention_mask: Optional[Union[torch.LongTensor, torch.BoolTensor]] = None,
**kwargs,
) -> Tuple[torch.FloatTensor, torch.FloatTensor]:
if attention_mask is not None:
attention_mask = attention_mask.bool()
else:
attention_mask = None
# MHA
if self.n_head == self.n_head_kv:
if past_key_values is None:
# If `past_key_values` are not supplied, we run self-attention
attn_output = self._forward_self_attn(x, attention_mask)
else:
# If `past_key_values` are supplied, it means that we might have cached values and
# could take advantage of cross-attention
attn_output = self._forward_cross_attn(
x, past_key_values, attention_mask
)
# MQA / GQA
else:
# Regardless of `past_key_values` being supplied or not, it always use cross-attention
# because `q` and `kv` lengths might be different
attn_output = self._forward_cross_attn(x, past_key_values, attention_mask)
output = rearrange(attn_output, "... h d -> ... (h d)")
output = self.out_proj(output)
return output if not self.return_residual else (output, x)
class ParallelBlock(nn.Module):
"""Parallel block.
This block applies parallel mixer and MLP layers to the input (used in GPT-J and CodeGen).
"""
def __init__(
self,
config: PretrainedConfig,
block_idx: Optional[int] = None,
) -> None:
super().__init__()
self.ln = nn.LayerNorm(config.n_embd, eps=config.layer_norm_epsilon)
self.resid_dropout = nn.Dropout(config.resid_pdrop)
self.block_idx = block_idx
self.mixer = MHA(config, layer_idx=block_idx)
self.mlp = MLP(config)
def forward(
self,
hidden_states: torch.FloatTensor,
past_key_values: Optional[Union[torch.FloatTensor, InferenceParams]] = None,
attention_mask: Optional[torch.BoolTensor] = None,
**kwargs,
) -> torch.FloatTensor:
residual = hidden_states
hidden_states = self.ln(hidden_states)
attn_outputs = self.mixer(
hidden_states,
past_key_values=past_key_values,
attention_mask=attention_mask,
)
if isinstance(attn_outputs, tuple):
attn_outputs = attn_outputs[0]
attn_outputs = self.resid_dropout(attn_outputs)
feed_forward_hidden_states = self.resid_dropout(self.mlp(hidden_states))
hidden_states = attn_outputs + feed_forward_hidden_states + residual
return hidden_states
class CausalLMHead(nn.Module):
"""Causal Language Modeling head.
Reference:
Improving Language Understanding by Generative Pre-Training.
https://cdn.openai.com/research-covers/language-unsupervised/language_understanding_paper.pdf.
"""
def __init__(self, config: PretrainedConfig) -> None:
super().__init__()
self.ln = nn.LayerNorm(config.n_embd, eps=config.layer_norm_epsilon)
self.linear = nn.Linear(config.n_embd, config.vocab_size)
def forward(self, hidden_states: torch.FloatTensor) -> torch.FloatTensor:
hidden_states = self.ln(hidden_states)
logits = self.linear(hidden_states).to(torch.float32)
return logits
class CausalLMLoss(nn.Module):
"""Causal Language Modeling loss.
Reference:
Improving Language Understanding by Generative Pre-Training.
https://cdn.openai.com/research-covers/language-unsupervised/language_understanding_paper.pdf.
"""
def __init__(self, shift_labels: bool = True) -> None:
super().__init__()
self.shift_labels = shift_labels
self.loss_fct = nn.CrossEntropyLoss()
def forward(
self, logits: torch.FloatTensor, labels: torch.LongTensor
) -> torch.FloatTensor:
if self.shift_labels:
logits = logits[..., :-1, :].contiguous()
labels = labels[..., 1:].contiguous()
loss = self.loss_fct(logits.view(-1, logits.size(-1)), labels.view(-1))
return loss
class PhiPreTrainedModel(PreTrainedModel):
"""Phi pre-trained model."""
config_class = PhiConfig
base_model_prefix = "transformer"
supports_gradient_checkpointing = False
_no_split_modules = ["ParallelBlock"]
def __init__(self, *inputs, **kwargs) -> None:
super().__init__(*inputs, **kwargs)
def prepare_inputs_for_generation(
self,
input_ids: torch.LongTensor = None,
inputs_embeds: torch.FloatTensor = None,
past_key_values: Optional[Union[torch.FloatTensor, InferenceParams]] = None,
attention_mask: Optional[Union[torch.LongTensor, torch.BoolTensor]] = None,
**kwargs,
) -> Dict[str, Any]:
if inputs_embeds is not None:
max_batch_size = inputs_embeds.shape[0]
seqlen_offset = inputs_embeds.shape[1] + input_ids.shape[1] - 2
elif input_ids is not None:
max_batch_size = input_ids.shape[0]
seqlen_offset = input_ids.shape[1] - 1
else:
raise ValueError(
"You have to specify either `input_ids` or `inputs_embeds`."
)
args = {}
if past_key_values is None or not (
isinstance(past_key_values, InferenceParams)
):
past_key_values = InferenceParams(
max_seqlen=self.config.n_positions,
max_batch_size=max_batch_size,
seqlen_offset=0,
batch_size_offset=0,
key_value_memory_dict={},
lengths_per_sample=None,
)
if inputs_embeds is not None:
args = {"inputs_embeds": inputs_embeds}
elif input_ids is not None:
args = {"input_ids": input_ids}
else:
raise ValueError(
"You have to specify either `input_ids` or `inputs_embeds`."
)
else:
# Assume that `past_key_values` has cached all tokens up to the last token in `input_ids`
past_key_values.seqlen_offset = seqlen_offset
input_ids = input_ids[:, -1].unsqueeze(-1)
args = {"input_ids": input_ids}
return {
**args,
"past_key_values": past_key_values,
"attention_mask": attention_mask,
}
class PhiModel(PhiPreTrainedModel):
"""Phi model."""
_keys_to_ignore_on_load_missing = [""]
_keys_to_ignore_on_load_unexpected = [r"h\.\d+\.mlp.(fc_in|fc_out)\.(weight|bias)"]
def __init__(self, config: PhiConfig) -> None:
super().__init__(config)
self.embd = Embedding(config)
self.h = nn.ModuleList(
[ParallelBlock(config, block_idx=i) for i in range(config.n_layer)]
)
self.gradient_checkpointing = config.gradient_checkpointing
self.post_init()
def get_input_embeddings(self) -> nn.Embedding:
return self.embd.wte
def set_input_embeddings(self, new_embeddings: nn.Embedding) -> None:
self.embd.wte = new_embeddings
def forward(
self,
input_ids: torch.LongTensor = None,
inputs_embeds: torch.FloatTensor = None,
past_key_values: Optional[Union[torch.FloatTensor, InferenceParams]] = None,
attention_mask: Optional[torch.BoolTensor] = None,
) -> torch.FloatTensor:
if input_ids is not None and inputs_embeds is not None:
raise ValueError(
"You cannot specify both `input_ids` and `inputs_embeds` at the same time."
)
elif input_ids is None and inputs_embeds is None:
raise ValueError(
"You have to specify either `input_ids` or `inputs_embeds`."
)
elif input_ids is not None:
hidden_states = self.embd(input_ids)
else:
hidden_states = inputs_embeds
for layer in self.h:
if self.gradient_checkpointing:
hidden_states = torch.utils.checkpoint.checkpoint(
layer.__call__,
hidden_states,
past_key_values,
attention_mask,
use_reentrant=True,
)
else:
hidden_states = layer(
hidden_states,
past_key_values=past_key_values,
attention_mask=attention_mask,
)
return hidden_states
class PhiForCausalLM(PhiPreTrainedModel):
"""Phi for Causal Language Modeling."""
_keys_to_ignore_on_load_missing = [""]
_keys_to_ignore_on_load_unexpected = [
r"transformer\.h\.\d+\.mlp.(fc_in|fc_out)\.(weight|bias)"
]
def __init__(self, config: PhiConfig) -> None:
super().__init__(config)
self.transformer = PhiModel(config)
self.lm_head = CausalLMHead(config)
self.loss = CausalLMLoss()
self.post_init()
def get_output_embeddings(self) -> nn.Linear:
return self.lm_head.linear
def set_output_embeddings(self, new_embeddings: nn.Linear) -> None:
self.lm_head.linear = new_embeddings
def forward(
self,
input_ids: torch.LongTensor = None,
inputs_embeds: torch.FloatTensor = None,
past_key_values: Optional[Union[torch.FloatTensor, InferenceParams]] = None,
attention_mask: Optional[torch.BoolTensor] = None,
labels: Optional[torch.LongTensor] = None,
**kwargs,
) -> CausalLMOutputWithPast:
hidden_states = self.transformer(
input_ids,
inputs_embeds,
past_key_values=past_key_values,
attention_mask=attention_mask,
)
lm_logits = self.lm_head(hidden_states)
loss = None
if labels is not None:
loss = self.loss(lm_logits, labels)
return CausalLMOutputWithPast(
loss=loss, logits=lm_logits, past_key_values=past_key_values
)
+10 -82
View File
@@ -1,91 +1,19 @@
import torch
from torch import nn
import transformers
from transformers import CodeGenTokenizerFast as Tokenizer
from accelerate import init_empty_weights, load_checkpoint_and_dispatch
from .phi.configuration_phi import PhiConfig
from .phi.modeling_phi import PhiForCausalLM
import re
from .modeling_phi import PhiForCausalLM
from .configuration_moondream import PhiConfig
transformers.logging.set_verbosity_error()
class TextModel:
def __init__(self, model_path: str = "model") -> None:
class TextModel(nn.Module):
def __init__(self, config) -> None:
super().__init__()
# Determine if CUDA (GPU) is available and use it; otherwise, use CPU
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
self.tokenizer = Tokenizer.from_pretrained(f"{model_path}/tokenizer")
phi_config = PhiConfig.from_pretrained(f"{model_path}/text_model_cfg.json")
with init_empty_weights():
self.model = PhiForCausalLM(phi_config)
self.model = load_checkpoint_and_dispatch(
self.model,
f"{model_path}/text_model.pt",
device_map={"": self.device.type},
).half()
self.text_emb = self.model.get_input_embeddings().to(self.device)
def input_embeds(self, prompt, image_embeds):
embeds = []
def _add_toks(toks):
embeds.append(self.text_emb(toks))
def _tokenize(txt):
return self.tokenizer(
txt, return_tensors="pt", add_special_tokens=False
).input_ids.to(self.model.device)
# Add BOS token
_add_toks(
torch.tensor([[self.tokenizer.bos_token_id]], device=self.model.device)
)
if "<image>" not in prompt:
embeds.append(self.text_emb(_tokenize(prompt)))
if type(config.phi_config) == dict:
phi_config = PhiConfig(**config.phi_config)
else:
assert prompt.count("<image>") == 1
before, after = prompt.split("<image>")
embeds.append(self.text_emb(_tokenize(f"{before}<image>")))
embeds.append(image_embeds.to(self.model.device))
embeds.append(self.text_emb(_tokenize(f"</image>{after}")))
phi_config = config.phi_config
return torch.cat(embeds, dim=1)
def generate(
self, image_embeds, prompt, eos_text="Human:", max_new_tokens=128, **kwargs
):
eos_tokens = self.tokenizer(eos_text, add_special_tokens=False)[0].ids
generate_config = {
"eos_token_id": eos_tokens,
"bos_token_id": self.tokenizer.bos_token_id,
"pad_token_id": self.tokenizer.eos_token_id,
"max_new_tokens": max_new_tokens,
**kwargs,
}
with torch.no_grad():
inputs_embeds = self.input_embeds(prompt, image_embeds)
output_ids = self.model.generate(
inputs_embeds=inputs_embeds, **generate_config
)
return self.tokenizer.batch_decode(output_ids, skip_special_tokens=True)
def answer_question(self, image_embeds, question, **kwargs):
prompt = f"<image>\n\nQuestion: {question}\n\nAnswer:"
answer = self.generate(
image_embeds,
prompt,
eos_text="<END>",
max_new_tokens=128,
**kwargs,
)[0]
return re.sub("<$", "", re.sub("END$", "", answer)).strip()
self.model = PhiForCausalLM(phi_config)
self.text_emb = self.model.get_input_embeddings()
+13
View File
@@ -0,0 +1,13 @@
import torch
def detect_device():
"""
Detects the appropriate device to run on, and return the device and dtype.
"""
if torch.cuda.is_available():
return torch.device("cuda"), torch.float16
elif torch.backends.mps.is_available():
return torch.device("mps"), torch.float16
else:
return torch.device("cpu"), torch.float32
+124 -10
View File
@@ -1,4 +1,5 @@
import torch
from torch import nn
from PIL import Image
from einops import rearrange
from torchvision.transforms.v2 import (
@@ -9,28 +10,141 @@ from torchvision.transforms.v2 import (
ToDtype,
Normalize,
)
import timm
class VisionEncoder:
def __init__(self, model_path: str = "model") -> None:
# Determine if CUDA (GPU) is available and use it; otherwise, use CPU
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
self.model = torch.jit.load(f"{model_path}/vision.pt").to(self.device).to(dtype=torch.float16)
class VisualHolder(nn.Module):
def __init__(self, model):
super().__init__()
self.visual = model
def forward(self, x):
return self.visual(x)
class ModelHolder(nn.Module):
def __init__(self, model):
super().__init__()
self.model = model
def forward(self, x):
return self.model(x)
class LinearPatchEmbedding(nn.Module):
def __init__(self, conv):
super().__init__()
self.linear = nn.Linear(588, 1152)
self.linear.weight.data = conv.weight.data.view(1152, -1)
if conv.bias is not None:
self.linear.bias.data = conv.bias.data
def forward(self, x):
return self.linear(x)
class MLP(nn.Module):
def __init__(
self,
in_features: int,
hidden_features: int = None,
out_features: int = None,
act_layer: nn.Module = nn.GELU,
) -> None:
super().__init__()
out_features = out_features or in_features
hidden_features = hidden_features or in_features
self.fc1 = nn.Linear(in_features, hidden_features)
self.act = act_layer()
self.fc2 = nn.Linear(hidden_features, out_features)
torch.nn.init.kaiming_normal_(
self.fc1.weight, mode="fan_in", nonlinearity="relu"
)
torch.nn.init.kaiming_normal_(
self.fc2.weight, mode="fan_in", nonlinearity="relu"
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.fc1(x)
x = self.act(x)
x = self.fc2(x)
return x
class VisionProjection(nn.Module):
def __init__(self):
super().__init__()
image_embedding_dim = 1152
model_dim = 2048
hidden_dim = model_dim * 4
self.mlp1 = MLP(image_embedding_dim, hidden_dim, model_dim)
self.mlp2 = MLP(model_dim, hidden_dim, model_dim)
self.ln = nn.LayerNorm(model_dim)
@property
def device(self):
return self.mlp1.fc1.weight.device
def forward(self, x):
x = self.mlp1(x)
x = self.ln(x)
x = x + self.mlp2(x)
return x
class VisionTower(nn.Module):
def __init__(self):
super().__init__()
self.encoder = ModelHolder(
VisualHolder(timm.create_model("vit_so400m_patch14_siglip_384"))
)
self.encoder.model.visual.patch_embed = LinearPatchEmbedding(
self.encoder.model.visual.patch_embed.proj
)
self.encoder.model.visual.attn_pool = nn.Identity()
self.projection = VisionProjection()
def forward(self, x):
x = self.encoder(x)
x = self.projection(x)
return x
class VisionEncoder(nn.Module):
def __init__(self) -> None:
super().__init__()
self.model = VisionTower()
self.preprocess = Compose(
[
Resize(size=(384, 384), interpolation=InterpolationMode.BICUBIC),
Resize(size=(378, 378), interpolation=InterpolationMode.BICUBIC),
ToImage(),
ToDtype(torch.float16, scale=True),
ToDtype(torch.float32, scale=True),
Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]),
]
)
@property
def device(self):
return self.model.projection.mlp1.fc1.weight.device
@property
def dtype(self):
return self.model.projection.mlp1.fc1.weight.dtype
def __call__(self, image: Image) -> torch.Tensor:
with torch.no_grad():
image_vec = self.preprocess(image.convert("RGB")).unsqueeze(0).to(self.device)
image_vec = image_vec[:, :, :-6, :-6]
image_vec = (
self.preprocess(image.convert("RGB"))
.unsqueeze(0)
.to(self.device, dtype=self.dtype)
)
image_vec = rearrange(
image_vec, "b c (h p1) (w p2) -> b (h w) (c p1 p2)", p1=14, p2=14
)
return self.model(image_vec)
+56 -22
View File
@@ -1,33 +1,62 @@
from moondream import VisionEncoder, TextModel
from PIL import Image
from huggingface_hub import snapshot_download
import torch
import argparse
from PIL import Image
from moondream import Moondream, detect_device
from queue import Queue
from threading import Thread
from transformers import TextIteratorStreamer
from transformers import (
TextIteratorStreamer,
CodeGenTokenizerFast as Tokenizer,
)
import re
model_path = snapshot_download("vikhyatk/moondream1")
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--image", type=str, required=True)
parser.add_argument("--prompt", type=str, required=False)
parser.add_argument("--cpu", action="store_true")
args = parser.parse_args()
vision_encoder = VisionEncoder(model_path)
text_model = TextModel(model_path)
if args.cpu:
device = torch.device("cpu")
dtype = torch.float32
else:
device, dtype = detect_device()
if device != torch.device("cpu"):
print("Using device:", device)
print("If you run into issues, pass the `--cpu` flag to this script.")
print()
parser = argparse.ArgumentParser()
parser.add_argument("--image", type=str, required=True)
parser.add_argument("--prompt", type=str, required=False)
args = parser.parse_args()
image_path = args.image
prompt = args.prompt
image = Image.open(args.image)
image_embeds = vision_encoder(image)
model_id = "vikhyatk/moondream1"
tokenizer = Tokenizer.from_pretrained(model_id)
moondream = Moondream.from_pretrained(model_id).to(device=device, dtype=dtype)
moondream.eval()
image = Image.open(image_path)
image_embeds = moondream.encode_image(image)
if prompt is None:
chat_history = ""
if args.prompt is None:
while True:
question = input("> ")
streamer = TextIteratorStreamer(text_model.tokenizer, skip_special_tokens=True)
generation_kwargs = dict(
image_embeds=image_embeds, question=question, streamer=streamer
result_queue = Queue()
streamer = TextIteratorStreamer(tokenizer, skip_special_tokens=True)
# Separate direct arguments from keyword arguments
thread_args = (image_embeds, question, tokenizer, chat_history)
thread_kwargs = {"streamer": streamer, "result_queue": result_queue}
thread = Thread(
target=moondream.answer_question,
args=thread_args,
kwargs=thread_kwargs,
)
thread = Thread(target=text_model.answer_question, kwargs=generation_kwargs)
thread.start()
buffer = ""
@@ -37,7 +66,12 @@ if args.prompt is None:
print(buffer, end="", flush=True)
buffer = ""
print(re.sub("<$", "", re.sub("END$", "", buffer)))
else:
question = args.prompt
print(">", question)
print(text_model.answer_question(image_embeds, question))
thread.join()
answer = result_queue.get()
chat_history += f"Question: {question}\n\nAnswer: {answer}\n\n"
else:
print(">", prompt)
answer = moondream.answer_question(image_embeds, prompt, tokenizer)
print(answer)