from dataclasses import dataclass import math from typing import Optional import torch import torch.nn as nn from torch.nn import functional as F from .base import ( Block, NormFunc, configure_optimizers, estimate_mfu, get_init_fn, init_weights_simple, ) TIE_WEIGHTS = False USE_SIN_POS_EMB = False SIMPLE_INIT = True def create_sin_embedding( positions: torch.Tensor, dim: int, max_period: float = 10_000, dtype: torch.dtype = torch.float32, ) -> torch.Tensor: """Create sinusoidal positional embedding, with shape `[B, T, C]`""" assert dim % 2 == 0 half_dim = dim // 2 positions = positions.to(dtype) adim = torch.arange(half_dim, device=positions.device, dtype=dtype).view(1, 1, -1) max_period_tensor = torch.full( [], max_period, device=positions.device, dtype=dtype ) # avoid sync point phase = positions / (max_period_tensor ** (adim / (half_dim - 1))) return torch.cat([torch.cos(phase), torch.sin(phase)], dim=-1) @dataclass class FineConfig: n_layer: int = 24 n_head: int = 16 # query heads n_kv_head: Optional[int] = None d_head: int = 64 bias: bool = False dropout: float = 0.0 coarse_vocab_size: int = 4160 coarse_codebook_size: int = 4096 coarse_n_codebooks: int = 8 coarse_pad_token: int = 4096 coarse_infer_token: int = 4097 coarse_rate_hz: int = 25 coarse_shift_factor: int = 0 coarse_samples: int = 50 # strided mask coarse_mask_period: int = 1 coarse_masked_samples: int = 0 fine_vocab_size: int = 1152 fine_codebook_size: int = 1024 fine_n_codebooks: int = 16 fine_pad_token: int = 1024 fine_infer_token: int = 1025 fine_rate_hz: int = 100 fine_shift_factor: int = 1 fine_samples: int = 200 t_memmap: int = 3375 # not sure what this is yet def __post_init__(self): # default to multi head attention if self.n_kv_head is None: self.n_kv_head = self.n_head assert self.coarse_masked_samples <= self.coarse_mask_period @property def n_embd(self): """The width of the residual stream""" return self.n_head * self.d_head @property def t_fine(self): """Includes infer token""" return self.fine_samples + self.fine_shift_factor * (self.fine_n_codebooks - 1) + 1 @property def t_coarse(self): return self.coarse_samples @property def block_size(self): return self.t_coarse + self.t_fine - 1 class Fine(nn.Module): def __init__(self, config: FineConfig): super().__init__() self.config = config model_dict = dict( wte_fine=nn.ModuleList( [ nn.Embedding(config.fine_vocab_size, config.n_embd) for _ in range(config.fine_n_codebooks) ] ), ln_fine=NormFunc(config.n_embd), wte_coarse=nn.ModuleList( [ nn.Embedding(config.coarse_vocab_size, config.n_embd) for _ in range(config.coarse_n_codebooks) ] ), ln_coarse=NormFunc(config.n_embd), drop=nn.Dropout(config.dropout), h=nn.ModuleList([Block(config) for _ in range(config.n_layer)]), ln_f=NormFunc(config.n_embd), ) if not USE_SIN_POS_EMB: model_dict["wpe"] = nn.Embedding(config.block_size, config.n_embd) self.transformer = nn.ModuleDict(model_dict) self.lm_heads = nn.ModuleList( [ nn.Linear(config.n_embd, config.fine_vocab_size, bias=False) for _ in range(config.fine_n_codebooks) ] ) if TIE_WEIGHTS: for n in range(config.fine_n_codebooks): self.transformer.wte_fine[n].weight = self.lm_heads[n].weight # init all weights if SIMPLE_INIT: self.apply(self._init_weights_simple) for pn, p in self.named_parameters(): if pn.endswith("c_proj.weight"): torch.nn.init.normal_(p, mean=0.0, std=0.02 / math.sqrt(2 * config.n_layer)) else: self._init_weights() print(f"number of parameters: {self.get_num_params()/1e6:.0f}M") def forward(self, x, y=None, coarse_offset=0, return_logits=False, last_only=True): device = x.device b, ns, t = x.size() assert ns == self.config.fine_n_codebooks + self.config.coarse_n_codebooks if y is not None: assert t == self.config.block_size _, _, t2 = y.size() assert t2 == self.config.t_fine - 1, (t2, self.config.t_fine) x_emb = 0 # embed coarse for n in range(self.config.coarse_n_codebooks): x_emb += self.transformer.ln_coarse(self.transformer.wte_coarse[n](x[:, n, :])) # embed fine for n in range(self.config.fine_n_codebooks): n2 = n + self.config.coarse_n_codebooks x_emb += self.transformer.ln_fine(self.transformer.wte_fine[n](x[:, n2, :])) # x_emb (b, t, n_embd) if USE_SIN_POS_EMB: pos = torch.arange(t, device=x.device).view(1, -1, 1) # pos = pos + offsets.view(-1, 1, 1) pos_emb = create_sin_embedding(pos, self.config.n_embd, dtype=x_emb.dtype) else: pos = torch.arange(t, dtype=torch.long, device=device).unsqueeze(0) # shape (1, t) pos_emb = self.transformer.wpe(pos) # (1, t, n_embd) x = self.transformer.drop(x_emb + pos_emb) for block in self.transformer.h: x = block(x) x = self.transformer.ln_f(x) x = x[:, coarse_offset : coarse_offset + self.config.t_fine - 1, :] if return_logits: if last_only: x = x[:, -1, :] fine_logits_list = [] for n in range(self.config.fine_n_codebooks): fine_logits_list.append(self.lm_heads[n](x)) fine_logits = torch.stack(fine_logits_list).swapaxes(0, 1) return fine_logits loss_dict = {} for n in range(self.config.fine_n_codebooks): logits = self.lm_heads[n](x) loss_dict[f"fine_{n}"] = F.cross_entropy( logits.reshape(-1, logits.size(-1)), y[:, n, :].reshape(-1), ignore_index=-1, ) return loss_dict def get_num_params(self, non_embedding=True): n_params = sum(p.numel() for p in self.parameters()) if non_embedding: for m in self.transformer.wte_coarse: n_params -= m.weight.numel() for m in self.transformer.wte_fine: n_params -= m.weight.numel() if not USE_SIN_POS_EMB: n_params -= self.transformer.wpe.weight.numel() return n_params def _init_weights_simple(self, module): init_weights_simple(self, module) def _init_weights(self): # embeddings get_init_fn(self.config.n_embd, init_depth=None)(self.transformer.wte_text.weight) for module in self.transformer.wte_coarse: get_init_fn(self.config.n_embd, init_depth=None)(module.weight) for module in self.transformer.wte_fine: get_init_fn(self.config.n_embd, init_depth=None)(module.weight) if not USE_SIN_POS_EMB: get_init_fn(self.config.n_embd, init_depth=None)(self.transformer.wpe.weight) # heads for module in self.lm_heads: get_init_fn(self.config.n_embd, init_depth=None)(module.weight) if module.bias is not None: torch.nn.init.zeros_(module.bias) # attention blocks for layer_idx, block in enumerate(self.transformer.h): # mlp module = block.mlp.c_fc get_init_fn(self.config.n_embd, init_depth=layer_idx + 1)(module.weight) if module.bias is not None: torch.nn.init.zeros_(module.bias) module = block.mlp.c_proj get_init_fn(block.mlp.embd_inner, init_depth=layer_idx + 1)(module.weight) if module.bias is not None: torch.nn.init.zeros_(module.bias) # attention for module in [block.attn.c_attn, block.attn.c_proj]: get_init_fn(self.config.n_embd, init_depth=layer_idx + 1)(module.weight) if module.bias is not None: torch.nn.init.zeros_(module.bias) def configure_optimizers(self, weight_decay, learning_rate, betas, device_type): return configure_optimizers(self, weight_decay, learning_rate, betas, device_type) def estimate_mfu(self, fwdbwd_per_iter, dt): return estimate_mfu(self, fwdbwd_per_iter, dt)