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