fela / modeling_fela.py
pandeyps's picture
Fela 1.6M
bd4f849 verified
Raw
History Blame Contribute Delete
9.88 kB
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from transformers import PreTrainedModel, PretrainedConfig
from transformers.modeling_outputs import CausalLMOutput
class Sin(nn.Module):
def __init__(self, dim, w=10, train_freq=True):
super().__init__()
self.freq = nn.Parameter(w * torch.ones(1, dim)) if train_freq else w * torch.ones(1, dim)
def forward(self, x):
return torch.sin(self.freq * x)
class PositionalEmbedding(nn.Module):
def __init__(self, emb_dim, seq_len):
super().__init__()
t = torch.linspace(0, 1, seq_len)[None, :, None]
bands = (emb_dim - 1) // 2
t_rescaled = torch.linspace(0, seq_len - 1, seq_len)[None, :, None]
w = 2 * math.pi * t_rescaled / seq_len
f = torch.linspace(1e-4, bands - 1, bands)[None, None]
z = torch.exp(-1j * f * w)
z = torch.cat([t, z.real, z.imag], dim=-1)
self.register_buffer("z", z)
self.register_buffer("t", t)
def forward(self, L):
return self.z[:, :L], self.t[:, :L]
class ExponentialModulation(nn.Module):
def __init__(self, d_model, fast_decay_pct=0.3, slow_decay_pct=1.5, target=1e-2, shift=0.0):
super().__init__()
self.shift = shift
max_decay = math.log(target) / fast_decay_pct
min_decay = math.log(target) / slow_decay_pct
deltas = torch.linspace(min_decay, max_decay, d_model)[None, None]
self.register_buffer("deltas", deltas)
def forward(self, t, x):
return x * (torch.exp(-t * self.deltas.abs()) + self.shift)
class HyenaFilter(nn.Module):
def __init__(self, d_model=256, emb_dim=5, order=64, seq_len=514,
num_inner_mlps=2, w=10, modulate=True):
super().__init__()
self.modulate = modulate
self.bias = nn.Parameter(torch.randn(d_model))
act = Sin(dim=order, w=w)
self.pos_emb = PositionalEmbedding(emb_dim, seq_len)
self.implicit_filter = nn.Sequential(
nn.Linear(emb_dim, order), act,
*[m for _ in range(num_inner_mlps) for m in (nn.Linear(order, order), act)],
nn.Linear(order, d_model, bias=False),
)
self.modulation = ExponentialModulation(d_model)
def filter(self, L):
z, t = self.pos_emb(L)
h = self.implicit_filter(z)
if self.modulate:
h = self.modulation(t, h)
return h
class ShortConv(nn.Module):
def __init__(self, d_model=256, order=2, short_filter_order=3):
super().__init__()
total_width = d_model * (order + 1)
self.in_proj = nn.Linear(d_model, total_width)
self.conv = nn.Conv1d(total_width, total_width, short_filter_order,
groups=total_width, padding=short_filter_order - 1)
def forward(self, u):
u = self.in_proj(u)
u = u.transpose(1, 2)
u = self.conv(u)[..., :u.shape[-1]]
return u
def fft_conv(u, k, bias=None):
seqlen = u.shape[-1]
fft_size = 2 * seqlen
k_f = torch.fft.rfft(k, n=fft_size) / fft_size
if len(u.shape) > 3:
k_f = k_f.unsqueeze(1)
u_f = torch.fft.rfft(u.to(dtype=k.dtype), n=fft_size)
y = torch.fft.irfft(u_f * k_f, n=fft_size, norm="forward")[..., :seqlen]
return y + u * bias.unsqueeze(-1) if bias is not None else y
class HyenaOperator(nn.Module):
def __init__(self, d_model=256, l_max=514, order=2, filter_order=64,
short_filter_order=3, drop_rate=0.0):
super().__init__()
self.d_model, self.order, self.l_max = d_model, order, l_max
self.in_proj = nn.Linear(d_model, (order + 1) * d_model)
self.out_proj = nn.Linear(d_model, d_model)
total_width = d_model * (order + 1)
self.short_filter = nn.Conv1d(total_width, total_width, short_filter_order,
groups=total_width, padding=short_filter_order - 1)
self.filter_fn = HyenaFilter(d_model=d_model, order=filter_order, seq_len=l_max)
self.dropout = nn.Dropout(drop_rate)
def forward(self, u):
l_filter = min(u.size(-2), self.l_max)
u = rearrange(self.in_proj(u), "b l d -> b d l")
uc = self.short_filter(u)[..., :l_filter]
*x, v = uc.split(self.d_model, dim=1)
k = self.filter_fn.filter(l_filter)
k = rearrange(k, "c l (v o) -> c o v l", v=self.d_model, o=self.order - 1)
bias = rearrange(self.filter_fn.bias, "(v o) -> o v", o=self.order - 1)
for o, x_i in enumerate(reversed(x[1:])):
v = self.dropout(v * x_i)
v = fft_conv(v, k[o], bias[o])
return self.out_proj(rearrange(v * x[0], "b v l -> b l v"))
class _Block(nn.Module):
def __init__(self, d_model, d_inner, l_max, drop1_p, drop2_p, eps=1e-5, residual_in_fp32=True):
super().__init__()
self.drop1 = nn.Dropout(drop1_p)
self.norm1 = nn.LayerNorm(d_model, eps=eps)
self.mixer = HyenaOperator(d_model=d_model, l_max=l_max)
self.drop2 = nn.Dropout(drop2_p)
self.norm2 = nn.LayerNorm(d_model, eps=eps)
self.mlp = nn.Sequential(nn.Linear(d_model, d_inner),
nn.GELU(approximate="tanh"),
nn.Linear(d_inner, d_model))
self.residual_in_fp32 = residual_in_fp32
def forward(self, hidden, residual):
dropped = self.drop1(hidden)
residual = (dropped + residual) if residual is not None else dropped
hidden = self.mixer(self.norm1(residual.to(dtype=self.norm1.weight.dtype)))
if self.residual_in_fp32:
residual = residual.float()
dropped = self.drop2(hidden)
residual = (dropped + residual) if residual is not None else dropped
hidden = self.mlp(self.norm2(residual.to(dtype=self.norm2.weight.dtype)))
if self.residual_in_fp32:
residual = residual.float()
return hidden, residual
class Fela(nn.Module):
def __init__(self, d_model=256, n_layer=2, d_inner=1024, vocab_size=32, l_max=514,
embed_dropout=0.1, resid_dropout=0.0, eps=1e-5, residual_in_fp32=True):
super().__init__()
torch.manual_seed(2222)
self.embed = nn.Embedding(vocab_size, d_model)
self.blocks = nn.ModuleList(
_Block(d_model, d_inner, l_max,
drop1_p=embed_dropout if i == 0 else resid_dropout,
drop2_p=resid_dropout, eps=eps,
residual_in_fp32=residual_in_fp32)
for i in range(n_layer)
)
self.drop_f = nn.Dropout(resid_dropout)
self.ln_f = nn.LayerNorm(d_model, eps=eps)
self.lm_head = nn.Linear(d_model, vocab_size, bias=False)
self._init_weights(n_layer)
self.lm_head.weight = self.embed.weight
def _init_weights(self, n_layer):
for m in self.modules():
if isinstance(m, nn.Linear):
nn.init.normal_(m.weight, std=0.02)
if m.bias is not None:
nn.init.zeros_(m.bias)
elif isinstance(m, nn.Embedding):
nn.init.normal_(m.weight, std=0.02)
for name, p in self.named_parameters():
if name.endswith("out_proj.weight") or name.endswith("mlp.2.weight"):
nn.init.normal_(p, std=0.02 / math.sqrt(2 * n_layer))
def forward(self, input_ids):
hidden = self.embed(input_ids)
residual = None
for block in self.blocks:
hidden, residual = block(hidden, residual)
dropped = self.drop_f(hidden)
residual = (dropped + residual) if residual is not None else dropped
hidden = self.ln_f(residual.to(dtype=self.ln_f.weight.dtype))
return self.lm_head(hidden)
try:
from transformers.generation import GenerationMixin
except ImportError:
from transformers.generation_utils import GenerationMixin
class FelaConfig(PretrainedConfig):
model_type = "fela"
def __init__(self, d_model=256, n_layer=2, d_inner=1024, vocab_size=32, l_max=514,
embed_dropout=0.1, resid_dropout=0.0, eps=1e-5, residual_in_fp32=True, **kwargs):
super().__init__(**kwargs)
self.d_model, self.n_layer, self.d_inner = d_model, n_layer, d_inner
self.vocab_size, self.l_max = vocab_size, l_max
self.embed_dropout, self.resid_dropout = embed_dropout, resid_dropout
self.eps, self.residual_in_fp32 = eps, residual_in_fp32
self.pad_token_id = 0
self.eos_token_id = 22
class FelaPreTrainedModel(PreTrainedModel):
config_class = FelaConfig
base_model_prefix = "fela"
def _init_weights(self, module):
pass
class FelaForCausalLM(FelaPreTrainedModel, GenerationMixin):
config_class = FelaConfig
base_model_prefix = "fela"
def __init__(self, config):
super().__init__(config)
self.fela = Fela(
d_model=config.d_model, n_layer=config.n_layer,
d_inner=config.d_inner, vocab_size=config.vocab_size,
l_max=config.l_max, embed_dropout=config.embed_dropout,
resid_dropout=config.resid_dropout, eps=config.eps,
residual_in_fp32=config.residual_in_fp32,
)
def get_input_embeddings(self):
return self.fela.embed
def set_input_embeddings(self, v):
self.fela.embed = v
self.fela.lm_head.weight = v.weight
def prepare_inputs_for_generation(self, input_ids, **kwargs):
return {"input_ids": input_ids}
def forward(self, input_ids=None, labels=None, attention_mask=None, **kwargs):
logits = self.fela(input_ids)
loss = None
if labels is not None:
loss = F.cross_entropy(logits.reshape(-1, self.config.vocab_size), labels.reshape(-1))
return CausalLMOutput(logits=logits, loss=loss)