from dataclasses import dataclass import torch import torch.nn as nn import torch.nn.functional as F class RMSNorm(nn.Module): def __init__(self, dimension: int, epsilon: float = 1e-6): super().__init__() self.weight = nn.Parameter(torch.ones(dimension)) self.epsilon = epsilon def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: normalized = hidden_states * torch.rsqrt( hidden_states.pow(2).mean(dim=-1, keepdim=True) + self.epsilon ) return self.weight * normalized class LiquidStateMixer(nn.Module): def __init__(self, dimension: int, kernel_size: int, dropout: float): super().__init__() self.kernel_size = kernel_size self.input_norm = RMSNorm(dimension) self.causal_depthwise_convolution = nn.Conv1d( dimension, dimension, kernel_size, groups=dimension, bias=True, ) self.state_parameters = nn.Linear(dimension, 3 * dimension) self.base_decay_logits = nn.Parameter(torch.zeros(dimension)) self.output_projection = nn.Linear(dimension, dimension, bias=False) self.dropout = nn.Dropout(dropout) def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: normalized = self.input_norm(hidden_states) convolution_input = normalized.transpose(1, 2) convolution_input = F.pad(convolution_input, (self.kernel_size - 1, 0)) local_features = self.causal_depthwise_convolution(convolution_input).transpose(1, 2) candidate, decay_logits, output_gate = self.state_parameters(local_features).chunk(3, dim=-1) candidate = torch.tanh(candidate) decay = torch.sigmoid(decay_logits + self.base_decay_logits) output_gate = torch.sigmoid(output_gate) state = torch.zeros_like(candidate[:, 0]) mixed_steps = [] for step in range(candidate.size(1)): step_decay = decay[:, step] state = step_decay * state + (1.0 - step_decay) * candidate[:, step] mixed_steps.append(output_gate[:, step] * state) mixed = torch.stack(mixed_steps, dim=1) return hidden_states + self.dropout(self.output_projection(mixed)) class SwiGLUFeedForward(nn.Module): def __init__(self, dimension: int, hidden_dimension: int, dropout: float): super().__init__() self.input_norm = RMSNorm(dimension) self.gate_projection = nn.Linear(dimension, hidden_dimension, bias=False) self.value_projection = nn.Linear(dimension, hidden_dimension, bias=False) self.output_projection = nn.Linear(hidden_dimension, dimension, bias=False) self.dropout = nn.Dropout(dropout) def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: normalized = self.input_norm(hidden_states) activated = F.silu(self.gate_projection(normalized)) * self.value_projection(normalized) return hidden_states + self.dropout(self.output_projection(activated)) class LiquidBlock(nn.Module): def __init__( self, dimension: int, hidden_dimension: int, kernel_size: int, dropout: float, ): super().__init__() self.state_mixer = LiquidStateMixer(dimension, kernel_size, dropout) self.feed_forward = SwiGLUFeedForward(dimension, hidden_dimension, dropout) def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: return self.feed_forward(self.state_mixer(hidden_states)) @dataclass class LiquidModelConfig: vocab_size: int dimension: int = 192 layer_count: int = 4 feed_forward_hidden_dimension: int = 512 convolution_kernel_size: int = 5 dropout: float = 0.05 pad_token_id: int = 0 bos_token_id: int = 2 eos_token_id: int = 3 class TinyLiquidCausalLanguageModel(nn.Module): def __init__(self, config: LiquidModelConfig): super().__init__() self.config = config self.token_embedding = nn.Embedding( config.vocab_size, config.dimension, padding_idx=config.pad_token_id, ) self.blocks = nn.ModuleList( [ LiquidBlock( config.dimension, config.feed_forward_hidden_dimension, config.convolution_kernel_size, config.dropout, ) for _ in range(config.layer_count) ] ) self.final_norm = RMSNorm(config.dimension) self.language_model_head = nn.Linear(config.dimension, config.vocab_size, bias=False) self.language_model_head.weight = self.token_embedding.weight self.apply(self._initialize_weights) @staticmethod def _initialize_weights(module: nn.Module) -> None: if isinstance(module, (nn.Linear, nn.Embedding)): nn.init.normal_(module.weight, mean=0.0, std=0.02) if isinstance(module, nn.Linear) and module.bias is not None: nn.init.zeros_(module.bias) def forward(self, input_ids: torch.Tensor, labels=None): hidden_states = self.token_embedding(input_ids) for block in self.blocks: hidden_states = block(hidden_states) logits = self.language_model_head(self.final_norm(hidden_states)) loss = None if labels is not None: loss = F.cross_entropy( logits.reshape(-1, logits.size(-1)), labels.reshape(-1), ignore_index=-100, ) return {"loss": loss, "logits": logits}