Files
kijai-ComfyUI-moondream/moondream/modeling_phi.py
T
Kijai 8cfbc3e4aa 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
2024-02-01 12:47:00 +02:00

721 lines
25 KiB
Python

# 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
)