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)