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:
@@ -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**
|
||||
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
)
|
||||
@@ -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
|
||||
@@ -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
@@ -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()
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
|
||||
vision_encoder = VisionEncoder(model_path)
|
||||
text_model = TextModel(model_path)
|
||||
|
||||
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()
|
||||
|
||||
image = Image.open(args.image)
|
||||
image_embeds = vision_encoder(image)
|
||||
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()
|
||||
|
||||
image_path = args.image
|
||||
prompt = args.prompt
|
||||
|
||||
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)))
|
||||
|
||||
thread.join()
|
||||
|
||||
answer = result_queue.get()
|
||||
chat_history += f"Question: {question}\n\nAnswer: {answer}\n\n"
|
||||
else:
|
||||
question = args.prompt
|
||||
print(">", question)
|
||||
print(text_model.answer_question(image_embeds, question))
|
||||
print(">", prompt)
|
||||
answer = moondream.answer_question(image_embeds, prompt, tokenizer)
|
||||
print(answer)
|
||||
|
||||
Reference in New Issue
Block a user