| import copy |
| import math |
|
|
| import torch |
| from torch import nn |
| from torch.nn import functional as F |
| from transformers import PretrainedConfig, PreTrainedModel |
| from transformers.modeling_outputs import ( |
| BaseModelOutput, |
| CausalLMOutput, |
| MaskedLMOutput, |
| ) |
|
|
|
|
| class TOLMConfig(PretrainedConfig): |
| model_type = "tolm" |
|
|
| def __init__( |
| self, |
| vocab_size=16000, |
| max_seq_len=512, |
| max_position_embeddings=None, |
| hidden_size=256, |
| num_hidden_layers=4, |
| num_attention_heads=4, |
| intermediate_size=1024, |
| position_buckets=32, |
| dropout=0.1, |
| hidden_dropout_prob=None, |
| attention_dropout=0.1, |
| attention_probs_dropout_prob=None, |
| initializer_range=0.03952847075210474, |
| layer_norm_eps=1.0e-5, |
| lm_head_gelu_approximate="tanh", |
| shared_relative_embeddings=False, |
| feedforward_dropout_after_projection=False, |
| attention_output_dropout=False, |
| embedding_padding_idx=True, |
| value_gating=True, |
| residual_mixing=True, |
| pad_token_id=1, |
| bos_token_id=2, |
| eos_token_id=3, |
| mask_token_id=4, |
| absolute_positions=False, |
| use_rope=False, |
| use_alibi=False, |
| recurrent_steps=1, |
| num_experts=1, |
| experts_per_token=1, |
| expert_intermediate_size=None, |
| future_offsets=None, |
| state_mixer_kernel=0, |
| geometry_lexical_dim=0, |
| geometry_curvature=1.0, |
| cognitive_readout_layer=0, |
| cognitive_readout_weight=0.0, |
| direct_sum_dims=None, |
| direct_sum_heads=None, |
| direct_sum_intermediate_sizes=None, |
| lexical_residual_buckets=0, |
| lexical_residual_dim=0, |
| lexical_residual_scale=1.0, |
| structured_projection_dim=0, |
| **kwargs, |
| ): |
| super().__init__( |
| pad_token_id=pad_token_id, |
| bos_token_id=bos_token_id, |
| eos_token_id=eos_token_id, |
| mask_token_id=mask_token_id, |
| **kwargs, |
| ) |
| self.vocab_size = vocab_size |
| self.max_seq_len = max_position_embeddings or max_seq_len |
| self.max_position_embeddings = self.max_seq_len |
| self.hidden_size = hidden_size |
| self.num_hidden_layers = num_hidden_layers |
| self.num_attention_heads = num_attention_heads |
| self.intermediate_size = intermediate_size |
| self.position_buckets = position_buckets |
| self.dropout = ( |
| hidden_dropout_prob if hidden_dropout_prob is not None else dropout |
| ) |
| self.hidden_dropout_prob = self.dropout |
| self.attention_dropout = ( |
| attention_probs_dropout_prob |
| if attention_probs_dropout_prob is not None |
| else attention_dropout |
| ) |
| self.attention_probs_dropout_prob = self.attention_dropout |
| self.initializer_range = initializer_range |
| self.layer_norm_eps = layer_norm_eps |
| self.lm_head_gelu_approximate = lm_head_gelu_approximate |
| self.shared_relative_embeddings = shared_relative_embeddings |
| self.feedforward_dropout_after_projection = ( |
| feedforward_dropout_after_projection |
| ) |
| self.attention_output_dropout = attention_output_dropout |
| self.embedding_padding_idx = embedding_padding_idx |
| self.value_gating = value_gating |
| self.residual_mixing = residual_mixing |
| self.absolute_positions = absolute_positions |
| self.use_rope = use_rope |
| self.use_alibi = use_alibi |
| self.recurrent_steps = recurrent_steps |
| self.num_experts = num_experts |
| self.experts_per_token = experts_per_token |
| self.expert_intermediate_size = expert_intermediate_size |
| self.future_offsets = future_offsets or [] |
| self.state_mixer_kernel = state_mixer_kernel |
| self.geometry_lexical_dim = geometry_lexical_dim |
| self.geometry_curvature = geometry_curvature |
| self.cognitive_readout_layer = cognitive_readout_layer |
| self.cognitive_readout_weight = cognitive_readout_weight |
| self.direct_sum_dims = direct_sum_dims or [] |
| self.direct_sum_heads = direct_sum_heads or [] |
| self.direct_sum_intermediate_sizes = direct_sum_intermediate_sizes or [] |
| self.lexical_residual_buckets = lexical_residual_buckets |
| self.lexical_residual_dim = lexical_residual_dim |
| self.lexical_residual_scale = lexical_residual_scale |
| self.structured_projection_dim = structured_projection_dim |
|
|
|
|
| def _valid_tokens(input_ids, attention_mask): |
| if attention_mask is None: |
| return torch.ones_like(input_ids, dtype=torch.bool) |
| return attention_mask.to(torch.bool) |
|
|
|
|
| def _bidirectional_mask(valid): |
| return valid[:, None, None, :] & valid[:, None, :, None] |
|
|
|
|
| def _causal_mask(valid): |
| length = valid.size(1) |
| causal = torch.ones((length, length), dtype=torch.bool, device=valid.device).tril() |
| return _bidirectional_mask(valid) & causal[None, None, :, :] |
|
|
|
|
| class RotaryPositionEncoding(nn.Module): |
| def __init__(self, head_width, max_length, *, base=10_000.0): |
| super().__init__() |
| if head_width % 2: |
| raise ValueError("RoPE head width must be even") |
| inv_freq = 1.0 / ( |
| base ** (torch.arange(0, head_width, 2, dtype=torch.float32) / head_width) |
| ) |
| self.register_buffer("inv_freq", inv_freq, persistent=False) |
| frequencies = self._frequencies(max_length, inv_freq.device) |
| self.register_buffer("cos", frequencies.cos(), persistent=False) |
| self.register_buffer("sin", frequencies.sin(), persistent=False) |
|
|
| def _frequencies(self, length, device): |
| positions = torch.arange(length, dtype=torch.float32, device=device) |
| return torch.outer(positions, self.inv_freq.to(device=device)) |
|
|
| def _rotate(self, value): |
| length = value.size(-2) |
| if length > self.cos.size(0): |
| frequencies = self._frequencies(length, value.device) |
| self.cos = frequencies.cos() |
| self.sin = frequencies.sin() |
| even, odd = value[..., 0::2], value[..., 1::2] |
| cos = self.cos[:length].to(device=value.device, dtype=value.dtype) |
| sin = self.sin[:length].to(device=value.device, dtype=value.dtype) |
| rotated = torch.stack((even * cos - odd * sin, even * sin + odd * cos), dim=-1) |
| return rotated.flatten(-2) |
|
|
| def encode(self, query, key): |
| return self._rotate(query), self._rotate(key) |
|
|
|
|
| class RelativeLogBucketSelfAttention(nn.Module): |
| def __init__(self, config): |
| super().__init__() |
| self.n_heads = config.num_attention_heads |
| self.d_head = config.hidden_size // config.num_attention_heads |
| self.max_seq_len = config.max_seq_len |
| self.buckets = config.position_buckets |
| self.value_gating = config.value_gating |
| self.use_rope = config.use_rope |
| self.use_alibi = config.use_alibi |
| self.shared_relative_embeddings = config.shared_relative_embeddings |
| self.attention_output_dropout = config.attention_output_dropout |
| if self.use_rope and self.use_alibi: |
| raise ValueError("RoPE and ALiBi are mutually exclusive") |
| self.qk = nn.Linear(config.hidden_size, 2 * config.hidden_size) |
| self.value = nn.Linear( |
| config.hidden_size, |
| 2 * config.hidden_size if config.value_gating else config.hidden_size, |
| ) |
| self.out = nn.Linear(config.hidden_size, config.hidden_size) |
| self.dropout = nn.Dropout(config.attention_dropout) |
| self.relative_embedding = ( |
| None |
| if self.use_rope or self.use_alibi or self.shared_relative_embeddings |
| else nn.Parameter( |
| torch.empty(2 * config.position_buckets - 1, config.hidden_size) |
| ) |
| ) |
| self.relative_norm = ( |
| None |
| if self.relative_embedding is None |
| else nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps) |
| ) |
| self.value_gate_norm = ( |
| nn.LayerNorm( |
| config.hidden_size, |
| eps=config.layer_norm_eps, |
| elementwise_affine=False, |
| ) |
| if config.value_gating |
| else None |
| ) |
| self.rope = ( |
| RotaryPositionEncoding(self.d_head, config.max_seq_len) |
| if config.use_rope |
| else None |
| ) |
| self.scale = 1.0 / math.sqrt( |
| self.d_head if config.use_rope or config.use_alibi else 3.0 * self.d_head |
| ) |
| if self.relative_embedding is not None: |
| nn.init.trunc_normal_( |
| self.relative_embedding, |
| mean=0.0, |
| std=config.initializer_range, |
| a=-2 * config.initializer_range, |
| b=2 * config.initializer_range, |
| ) |
| self.register_buffer( |
| "position_indices", |
| self._position_indices(config.max_seq_len, torch.device("cpu")), |
| persistent=False, |
| ) |
| self.register_buffer( |
| "alibi_bias", |
| self._alibi_bias(config.max_seq_len, torch.device("cpu")), |
| persistent=False, |
| ) |
|
|
| def _alibi_bias(self, length, device): |
| positions = torch.arange(length, device=device) |
| distance = (positions[:, None] - positions[None, :]).abs().float() |
| slopes = torch.pow( |
| 2.0, |
| -8.0 |
| * (torch.arange(self.n_heads, device=device).float() + 1.0) |
| / self.n_heads, |
| ) |
| return -slopes[None, :, None, None] * distance[None, None, :, :] |
|
|
| def _position_indices(self, length, device): |
| positions = torch.arange(length, device=device) |
| relative = positions[:, None] - positions[None, :] |
| sign = torch.sign(relative) |
| half = self.buckets // 2 |
| absolute = relative.abs().clamp(max=max(half + 1, self.max_seq_len - 1)) |
| near = absolute <= half |
| safe = absolute.clamp_min(half) |
| denominator = math.log(max((self.max_seq_len - 1) / half, 1.0001)) |
| logged = ( |
| torch.ceil(torch.log(safe / half) / denominator * (half - 1)).long() + half |
| ) |
| bucketed = torch.where(near, relative, logged * sign) |
| return ( |
| bucketed.long().clamp(-self.buckets + 1, self.buckets - 1) |
| + self.buckets |
| - 1 |
| ) |
|
|
| def forward(self, x, mask, relative_embedding=None): |
| batch, length, width = x.shape |
| if length > self.position_indices.size(0): |
| self.position_indices = self._position_indices(length, x.device) |
| q, k = self.qk(x).chunk(2, dim=-1) |
| if self.value_gating: |
| v, gate = self.value(x).chunk(2, dim=-1) |
| gate = F.gelu(gate) |
| else: |
| v, gate = self.value(x), None |
| q = q.view(batch, length, self.n_heads, self.d_head).transpose(1, 2) |
| k = k.view(batch, length, self.n_heads, self.d_head).transpose(1, 2) |
| v = v.view(batch, length, self.n_heads, self.d_head).transpose(1, 2) |
|
|
| if self.rope is not None: |
| q, k = self.rope.encode(q, k) |
| scores = torch.matmul(q, k.transpose(-2, -1)) * self.scale |
| elif self.use_alibi: |
| if length > self.alibi_bias.size(-1): |
| self.alibi_bias = self._alibi_bias(length, x.device) |
| scores = torch.matmul(q, k.transpose(-2, -1)) * self.scale |
| scores = scores + self.alibi_bias[:, :, :length, :length].to( |
| device=x.device, dtype=scores.dtype |
| ) |
| else: |
| if relative_embedding is None: |
| assert self.relative_embedding is not None |
| assert self.relative_norm is not None |
| relative_embedding = self.relative_norm(self.relative_embedding) |
| relative = self.qk(self.dropout(relative_embedding)) |
| relative = relative[self.position_indices[:length, :length].to(x.device)] |
| q_pos, k_pos = relative.chunk(2, dim=-1) |
| q_pos = q_pos.view(length, length, self.n_heads, self.d_head) |
| k_pos = k_pos.view(length, length, self.n_heads, self.d_head) |
| scores = torch.matmul(q, k.transpose(-2, -1)) * self.scale |
| scores = scores + torch.einsum("bhqd,qkhd->bhqk", q, k_pos) * self.scale |
| scores = scores + torch.einsum("bhkd,qkhd->bhqk", k, q_pos) * self.scale |
|
|
| probs = torch.softmax( |
| scores.masked_fill(~mask, torch.finfo(scores.dtype).min), dim=-1 |
| ) |
| probs = self.dropout(probs) if self.training else probs |
| output = ( |
| torch.matmul(probs, v) |
| .transpose(1, 2) |
| .contiguous() |
| .view(batch, length, width) |
| ) |
| if gate is not None and self.value_gate_norm is not None: |
| output = self.value_gate_norm(output * gate) |
| output = self.out(output) |
| return self.dropout(output) if self.attention_output_dropout else output |
|
|
|
|
| class GeGLU(nn.Module): |
| def __init__(self, config, width=None): |
| super().__init__() |
| width = width or config.intermediate_size |
| self.up = nn.Linear(config.hidden_size, 2 * width, bias=False) |
| self.post_activation_norm = nn.LayerNorm( |
| width, eps=config.layer_norm_eps, elementwise_affine=False |
| ) |
| self.down = nn.Linear(width, config.hidden_size, bias=False) |
| self.dropout = nn.Dropout(config.dropout) |
| self.dropout_after_projection = config.feedforward_dropout_after_projection |
|
|
| def forward(self, x): |
| value, gate = self.up(x).chunk(2, dim=-1) |
| hidden = value * F.gelu(gate, approximate="tanh") |
| hidden = self.post_activation_norm(hidden) |
| if self.dropout_after_projection: |
| return self.dropout(self.down(hidden)) |
| return self.down(self.dropout(hidden)) |
|
|
|
|
| class RoutedGeGLU(nn.Module): |
| def __init__(self, config): |
| super().__init__() |
| if not 1 <= config.experts_per_token <= config.num_experts: |
| raise ValueError("experts_per_token must be in [1, num_experts]") |
| width = config.expert_intermediate_size or max( |
| 1, config.intermediate_size // config.num_experts |
| ) |
| self.top_k = config.experts_per_token |
| self.router = nn.Linear(config.hidden_size, config.num_experts, bias=False) |
| self.experts = nn.ModuleList( |
| GeGLU(config, width) for _ in range(config.num_experts) |
| ) |
|
|
| def forward(self, x): |
| probabilities = self.router(x).softmax(dim=-1) |
| weights, indices = probabilities.topk(self.top_k, dim=-1) |
| weights = weights / weights.sum(dim=-1, keepdim=True).clamp_min(1e-8) |
| gates = torch.zeros_like(probabilities).scatter(-1, indices, weights) |
| outputs = torch.stack([expert(x) for expert in self.experts], dim=-2) |
| return (outputs * gates.unsqueeze(-1)).sum(dim=-2) |
|
|
|
|
| class CausalStateMixer(nn.Module): |
| def __init__(self, config): |
| super().__init__() |
| kernel = int(config.state_mixer_kernel) |
| self.input = nn.Linear(config.hidden_size, 2 * config.hidden_size, bias=False) |
| self.state = nn.Conv1d( |
| config.hidden_size, |
| config.hidden_size, |
| kernel, |
| groups=config.hidden_size, |
| padding=kernel - 1, |
| bias=False, |
| ) |
| self.output = nn.Linear(config.hidden_size, config.hidden_size, bias=False) |
| self.gate = nn.Parameter(torch.tensor(-2.0)) |
|
|
| def forward(self, hidden): |
| value, gate = self.input(hidden).chunk(2, dim=-1) |
| state = self.state(value.transpose(1, 2))[..., : hidden.size(1)].transpose(1, 2) |
| return self.output(state * F.silu(gate)) * self.gate.sigmoid() |
|
|
|
|
| class DynamicWeightedAverage(nn.Module): |
| def __init__(self, n_sublayers): |
| super().__init__() |
| self.alphas = nn.ParameterList( |
| nn.Parameter(torch.cat([torch.zeros(i + 1), torch.ones(1)])) |
| for i in range(int(n_sublayers)) |
| ) |
| self._states = None |
|
|
| def initialize(self, hidden): |
| self._states = [hidden] |
|
|
| def forward(self, hidden, sublayer_index): |
| self._states.append(hidden) |
| return torch.tensordot( |
| self.alphas[sublayer_index], torch.stack(self._states), dims=1 |
| ) |
|
|
|
|
| class GPTBertBlock(nn.Module): |
| def __init__(self, config): |
| super().__init__() |
| self.attention_norm = nn.LayerNorm( |
| config.hidden_size, |
| eps=config.layer_norm_eps, |
| elementwise_affine=False, |
| ) |
| self.attention = RelativeLogBucketSelfAttention(config) |
| self.state_mixer = ( |
| CausalStateMixer(config) if config.state_mixer_kernel else None |
| ) |
| self.feedforward_norm = nn.LayerNorm( |
| config.hidden_size, |
| eps=config.layer_norm_eps, |
| elementwise_affine=False, |
| ) |
| self.feedforward = ( |
| RoutedGeGLU(config) if config.num_experts > 1 else GeGLU(config) |
| ) |
|
|
| def attend(self, hidden, mask, relative_embedding=None): |
| normalized = self.attention_norm(hidden) |
| attention = self.attention(normalized, mask, relative_embedding) |
| return ( |
| attention |
| if self.state_mixer is None |
| else attention + self.state_mixer(normalized) |
| ) |
|
|
| def transform(self, hidden): |
| return self.feedforward(self.feedforward_norm(hidden)) |
|
|
|
|
| class LexicalResidualEmbedding(nn.Module): |
| def __init__(self, config): |
| super().__init__() |
| buckets = int(config.lexical_residual_buckets) |
| width = int(config.lexical_residual_dim) or config.hidden_size |
| self.buckets = buckets |
| self.pad_token_id = config.pad_token_id |
| self.scale = float(config.lexical_residual_scale) |
| self.embedding = nn.Embedding(buckets, width, padding_idx=0) |
| self.projection = ( |
| nn.Identity() |
| if width == config.hidden_size |
| else nn.Linear(width, config.hidden_size, bias=False) |
| ) |
| self.register_buffer( |
| "word_start_vocab_mask", |
| torch.zeros(config.vocab_size, dtype=torch.bool), |
| ) |
|
|
| def forward(self, token_ids): |
| batch, length = token_ids.shape |
| positions = torch.arange(length, device=token_ids.device).expand(batch, -1) |
| starts = self.word_start_vocab_mask[token_ids].clone() |
| starts[:, 0] = True |
| start_positions = torch.where(starts, positions, 0).cummax(dim=1).values |
| relative = positions - start_positions |
| ordinal = relative.long() + 1 |
| mixed = (token_ids.long() + 1) * ordinal * 1_000_003 + ordinal * 97_409 |
| cumulative = mixed.cumsum(dim=1) |
| before = F.pad(cumulative[:, :-1], (1, 0)) |
| prefix_hashes = cumulative - before.gather(1, start_positions) |
| buckets = prefix_hashes.remainder(self.buckets - 1) + 1 |
| buckets = buckets.masked_fill(token_ids.eq(self.pad_token_id), 0) |
| return self.scale * self.projection(self.embedding(buckets)) |
|
|
|
|
| class GPTBertBackbone(nn.Module): |
| def __init__(self, config): |
| super().__init__() |
| self.config = config |
| self.embed_tokens = nn.Embedding( |
| config.vocab_size, |
| config.hidden_size, |
| padding_idx=config.pad_token_id if config.embedding_padding_idx else None, |
| ) |
| self.relative_embedding = ( |
| nn.Parameter( |
| torch.empty( |
| 2 * config.position_buckets - 1, config.hidden_size |
| ) |
| ) |
| if config.shared_relative_embeddings |
| else None |
| ) |
| self.relative_norm = ( |
| nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps) |
| if config.shared_relative_embeddings |
| else None |
| ) |
| if self.relative_embedding is not None: |
| nn.init.trunc_normal_( |
| self.relative_embedding, |
| std=config.initializer_range, |
| a=-2 * config.initializer_range, |
| b=2 * config.initializer_range, |
| ) |
| self.geometry_lexical_dim = int(config.geometry_lexical_dim) |
| if not 0 <= self.geometry_lexical_dim < config.hidden_size: |
| raise ValueError("geometry_lexical_dim must be in [0, hidden_size)") |
| self.geometry_curvature = float(config.geometry_curvature) |
| if self.geometry_curvature <= 0: |
| raise ValueError("geometry_curvature must be positive") |
| self.lexical_angle = None |
| self.lexical_radius = None |
| if self.geometry_lexical_dim: |
| self.lexical_angle = nn.Embedding( |
| config.vocab_size, self.geometry_lexical_dim, config.pad_token_id |
| ) |
| self.lexical_radius = nn.Embedding(config.vocab_size, 1, config.pad_token_id) |
| self.lexical_residual = ( |
| LexicalResidualEmbedding(config) |
| if config.lexical_residual_buckets |
| else None |
| ) |
| self.embed_positions = ( |
| nn.Embedding(config.max_seq_len, config.hidden_size) |
| if getattr(config, "absolute_positions", False) |
| else None |
| ) |
| self.embed_norm = nn.LayerNorm( |
| config.hidden_size, |
| eps=config.layer_norm_eps, |
| elementwise_affine=False, |
| ) |
| self.dropout = nn.Dropout(config.dropout) |
| self.blocks = nn.ModuleList( |
| GPTBertBlock(config) for _ in range(config.num_hidden_layers) |
| ) |
| self.recurrent_steps = max(1, int(config.recurrent_steps)) |
| self.future_projections = nn.ModuleDict( |
| { |
| str(offset): nn.Linear( |
| config.hidden_size, config.hidden_size, bias=False |
| ) |
| for offset in config.future_offsets |
| } |
| ) |
| self.residual_mixer = ( |
| DynamicWeightedAverage(config.num_hidden_layers * self.recurrent_steps * 2) |
| if config.residual_mixing |
| else None |
| ) |
| self.cognitive_readout_layer = int(config.cognitive_readout_layer) |
| self.cognitive_readout_weight = float(config.cognitive_readout_weight) |
| maximum_depth = config.num_hidden_layers * self.recurrent_steps |
| if self.cognitive_readout_layer > maximum_depth: |
| raise ValueError("cognitive_readout_layer exceeds the executed depth") |
|
|
| def lexical_geometry(self, token_ids): |
| if self.lexical_angle is None or self.lexical_radius is None: |
| raise RuntimeError("lexical geometry is disabled") |
| direction = F.normalize(self.lexical_angle(token_ids), dim=-1) |
| radius = F.softplus(self.lexical_radius(token_ids)).squeeze(-1) |
| scale = math.sqrt(self.geometry_curvature) |
| point = torch.tanh(scale * radius / 2).unsqueeze(-1) * direction / scale |
| return point, radius |
|
|
| def forward(self, input_ids, mask): |
| embedded = self.embed_tokens(input_ids) |
| if self.lexical_residual is not None: |
| embedded = embedded + self.lexical_residual(input_ids) |
| if self.geometry_lexical_dim: |
| _, radius = self.lexical_geometry(input_ids) |
| direction = F.normalize(self.lexical_angle(input_ids), dim=-1) |
| tangent = radius.unsqueeze(-1) * direction |
| embedded = torch.cat((embedded[..., :-self.geometry_lexical_dim], tangent), -1) |
| if self.embed_positions is not None: |
| positions = torch.arange(input_ids.size(1), device=input_ids.device) |
| embedded = embedded + self.embed_positions(positions) |
| hidden = self.dropout(self.embed_norm(embedded)) |
| mixer = self.residual_mixer |
| if mixer is not None: |
| mixer.initialize(hidden) |
| sublayer = 0 |
| cognitive_hidden = None |
| layer_index = 0 |
| relative = ( |
| self.relative_norm(self.relative_embedding) |
| if self.relative_norm is not None and self.relative_embedding is not None |
| else None |
| ) |
| for _ in range(self.recurrent_steps): |
| for block in self.blocks: |
| hidden = hidden + block.attend(hidden, mask, relative) |
| if mixer is not None: |
| hidden = mixer(hidden, sublayer) |
| sublayer += 1 |
| hidden = hidden + block.transform(hidden) |
| if mixer is not None: |
| hidden = mixer(hidden, sublayer) |
| sublayer += 1 |
| layer_index += 1 |
| if layer_index == self.cognitive_readout_layer: |
| cognitive_hidden = hidden |
| if cognitive_hidden is not None and self.cognitive_readout_weight > 0: |
| weight = self.cognitive_readout_weight |
| hidden = (1.0 - weight) * hidden + weight * cognitive_hidden |
| return hidden |
|
|
|
|
| class DirectSumStream(nn.Module): |
| def __init__(self, config): |
| super().__init__() |
| self.relative_embedding = ( |
| nn.Parameter(torch.empty(2 * config.position_buckets - 1, config.hidden_size)) |
| if config.shared_relative_embeddings |
| else None |
| ) |
| self.relative_norm = ( |
| nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps) |
| if config.shared_relative_embeddings |
| else None |
| ) |
| if self.relative_embedding is not None: |
| nn.init.trunc_normal_( |
| self.relative_embedding, |
| std=config.initializer_range, |
| a=-2 * config.initializer_range, |
| b=2 * config.initializer_range, |
| ) |
| self.embed_positions = ( |
| nn.Embedding(config.max_seq_len, config.hidden_size) |
| if config.absolute_positions |
| else None |
| ) |
| self.embed_norm = nn.LayerNorm( |
| config.hidden_size, |
| eps=config.layer_norm_eps, |
| elementwise_affine=False, |
| ) |
| self.dropout = nn.Dropout(config.dropout) |
| self.blocks = nn.ModuleList( |
| GPTBertBlock(config) for _ in range(config.num_hidden_layers) |
| ) |
| self.recurrent_steps = max(1, int(config.recurrent_steps)) |
| self.residual_mixer = ( |
| DynamicWeightedAverage(config.num_hidden_layers * self.recurrent_steps * 2) |
| if config.residual_mixing |
| else None |
| ) |
|
|
| def forward(self, embedded, mask): |
| if self.embed_positions is not None: |
| positions = torch.arange(embedded.size(1), device=embedded.device) |
| embedded = embedded + self.embed_positions(positions) |
| hidden = self.dropout(self.embed_norm(embedded)) |
| mixer = self.residual_mixer |
| if mixer is not None: |
| mixer.initialize(hidden) |
| relative = ( |
| self.relative_norm(self.relative_embedding) |
| if self.relative_norm is not None and self.relative_embedding is not None |
| else None |
| ) |
| sublayer = 0 |
| for _ in range(self.recurrent_steps): |
| for block in self.blocks: |
| hidden = hidden + block.attend(hidden, mask, relative) |
| if mixer is not None: |
| hidden = mixer(hidden, sublayer) |
| sublayer += 1 |
| hidden = hidden + block.transform(hidden) |
| if mixer is not None: |
| hidden = mixer(hidden, sublayer) |
| sublayer += 1 |
| return hidden |
|
|
|
|
| class DirectSumBackbone(nn.Module): |
| def __init__(self, config): |
| super().__init__() |
| dims = tuple(int(value) for value in config.direct_sum_dims) |
| heads = tuple(int(value) for value in config.direct_sum_heads) |
| widths = tuple(int(value) for value in config.direct_sum_intermediate_sizes) |
| if len(dims) != 3 or len(heads) != 3 or len(widths) != 3: |
| raise ValueError("direct sum requires three dims, heads and FFN widths") |
| if sum(dims) != config.hidden_size: |
| raise ValueError("direct_sum_dims must sum to hidden_size") |
| if any(dim % head for dim, head in zip(dims, heads)): |
| raise ValueError("each direct-sum dimension must divide its head count") |
| self.dims = dims |
| self.embed_tokens = nn.Embedding( |
| config.vocab_size, |
| config.hidden_size, |
| padding_idx=config.pad_token_id if config.embedding_padding_idx else None, |
| ) |
| streams = [] |
| for dim, head, width in zip(dims, heads, widths): |
| stream_config = copy.copy(config) |
| stream_config.hidden_size = dim |
| stream_config.num_attention_heads = head |
| stream_config.intermediate_size = width |
| stream_config.direct_sum_dims = [] |
| stream_config.direct_sum_heads = [] |
| stream_config.direct_sum_intermediate_sizes = [] |
| stream_config.geometry_lexical_dim = 0 |
| stream_config.future_offsets = [] |
| stream_config.cognitive_readout_layer = 0 |
| stream_config.cognitive_readout_weight = 0.0 |
| streams.append(DirectSumStream(stream_config)) |
| self.streams = nn.ModuleList(streams) |
| self.concept_radius = nn.Embedding( |
| config.vocab_size, 1, padding_idx=config.pad_token_id |
| ) |
|
|
| @property |
| def factor_slices(self): |
| syntax, lexical, conceptual = self.dims |
| return ( |
| slice(0, syntax), |
| slice(syntax, syntax + lexical), |
| slice(syntax + lexical, syntax + lexical + conceptual), |
| ) |
|
|
| def conceptual_geometry(self, token_ids): |
| conceptual = self.embed_tokens(token_ids)[..., self.factor_slices[2]] |
| direction = F.normalize(conceptual, dim=-1) |
| radius = (1.0 - 1.0e-4) * torch.sigmoid( |
| self.concept_radius(token_ids).squeeze(-1) |
| ) |
| if self.concept_radius.padding_idx is not None: |
| radius = radius.masked_fill( |
| token_ids.eq(self.concept_radius.padding_idx), 0.0 |
| ) |
| return radius.unsqueeze(-1) * direction, radius |
|
|
| def forward(self, input_ids, mask): |
| embedded = self.embed_tokens(input_ids) |
| conceptual, _ = self.conceptual_geometry(input_ids) |
| parts = list(embedded.split(self.dims, dim=-1)) |
| parts[2] = conceptual |
| return torch.cat( |
| [stream(part, mask) for stream, part in zip(self.streams, parts)], dim=-1 |
| ) |
|
|
|
|
| def _backbone(config): |
| return DirectSumBackbone(config) if config.direct_sum_dims else GPTBertBackbone(config) |
|
|
|
|
| def _structured_projections(config): |
| width = int(config.structured_projection_dim) |
| return ( |
| nn.ModuleDict( |
| { |
| "syntax": nn.Linear(config.hidden_size, width, bias=False), |
| "lexical": nn.Linear(config.hidden_size, width, bias=False), |
| } |
| ) |
| if width |
| else nn.ModuleDict() |
| ) |
|
|
|
|
| class GPTBertLMHead(nn.Module): |
| def __init__(self, config): |
| super().__init__() |
| self.norm = nn.LayerNorm( |
| config.hidden_size, |
| eps=config.layer_norm_eps, |
| elementwise_affine=False, |
| ) |
| self.dense = nn.Linear(config.hidden_size, config.hidden_size) |
| self.post_norm = nn.LayerNorm( |
| config.hidden_size, |
| eps=config.layer_norm_eps, |
| elementwise_affine=False, |
| ) |
| self.dropout = nn.Dropout(config.dropout) |
| self._approximate = config.lm_head_gelu_approximate |
| self.bias = nn.Parameter(torch.zeros(config.vocab_size)) |
|
|
| def forward(self, hidden): |
| projected = self.dropout( |
| self.post_norm( |
| F.gelu( |
| self.dense(self.norm(hidden)), approximate=self._approximate |
| ) |
| ) |
| ) |
| return F.linear(projected, self.weight, self.bias) |
|
|
|
|
| class TOLMModel(PreTrainedModel): |
| config_class = TOLMConfig |
| base_model_prefix = "tolm" |
| _no_split_modules = ["GPTBertBlock"] |
|
|
| def __init__(self, config): |
| super().__init__(config) |
| self.backbone = _backbone(config) |
| self.structured_projections = _structured_projections(config) |
| self.post_init() |
|
|
| def get_input_embeddings(self): |
| return self.backbone.embed_tokens |
|
|
| def forward(self, input_ids=None, attention_mask=None, **kwargs): |
| if input_ids is None: |
| raise ValueError("input_ids is required") |
| hidden = self.backbone( |
| input_ids, _bidirectional_mask(_valid_tokens(input_ids, attention_mask)) |
| ) |
| return BaseModelOutput( |
| last_hidden_state=hidden, hidden_states=None, attentions=None |
| ) |
|
|
|
|
| class TOLMForMaskedLM(PreTrainedModel): |
| config_class = TOLMConfig |
| base_model_prefix = "tolm" |
| _no_split_modules = ["GPTBertBlock"] |
| _tied_weights_keys = ["heads.lm.weight"] |
|
|
| def __init__(self, config): |
| super().__init__(config) |
| self.backbone = _backbone(config) |
| head = GPTBertLMHead(config) |
| self.heads = nn.ModuleDict({"lm": head}) |
| self.structured_projections = _structured_projections(config) |
| if config.direct_sum_dims: |
| self.factor_dual_lambdas = nn.Parameter( |
| torch.ones(3), requires_grad=False |
| ) |
| self.post_init() |
|
|
| def get_input_embeddings(self): |
| return self.backbone.embed_tokens |
|
|
| def get_output_embeddings(self): |
| return self.heads["lm"] |
|
|
| def forward(self, input_ids=None, attention_mask=None, labels=None, **kwargs): |
| if input_ids is None: |
| raise ValueError("input_ids is required") |
| hidden = self.backbone( |
| input_ids, _bidirectional_mask(_valid_tokens(input_ids, attention_mask)) |
| ) |
| logits = self.heads["lm"](hidden) |
| loss = ( |
| None |
| if labels is None |
| else F.cross_entropy( |
| logits.view(-1, logits.size(-1)), labels.view(-1), ignore_index=-100 |
| ) |
| ) |
| return MaskedLMOutput( |
| loss=loss, logits=logits, hidden_states=None, attentions=None |
| ) |
|
|
|
|
| class TOLMForCausalLM(PreTrainedModel): |
| config_class = TOLMConfig |
| base_model_prefix = "tolm" |
| _no_split_modules = ["GPTBertBlock"] |
| _tied_weights_keys = ["heads.lm.weight"] |
|
|
| def __init__(self, config): |
| super().__init__(config) |
| self.backbone = _backbone(config) |
| head = GPTBertLMHead(config) |
| self.heads = nn.ModuleDict({"lm": head}) |
| self.structured_projections = _structured_projections(config) |
| if config.direct_sum_dims: |
| self.factor_dual_lambdas = nn.Parameter( |
| torch.ones(3), requires_grad=False |
| ) |
| self.post_init() |
|
|
| def get_input_embeddings(self): |
| return self.backbone.embed_tokens |
|
|
| def get_output_embeddings(self): |
| return self.heads["lm"] |
|
|
| def prepare_inputs_for_generation(self, input_ids, attention_mask=None, **kwargs): |
| return {"input_ids": input_ids, "attention_mask": attention_mask} |
|
|
| def forward(self, input_ids=None, attention_mask=None, labels=None, **kwargs): |
| if input_ids is None: |
| raise ValueError("input_ids is required") |
| hidden = self.backbone( |
| input_ids, _causal_mask(_valid_tokens(input_ids, attention_mask)) |
| ) |
| logits = self.heads["lm"](hidden) |
| loss = None |
| if labels is not None: |
| loss = F.cross_entropy( |
| logits[:, :-1].contiguous().view(-1, logits.size(-1)), |
| labels[:, 1:].contiguous().view(-1), |
| ignore_index=-100, |
| ) |
| return CausalLMOutput( |
| loss=loss, logits=logits, hidden_states=None, attentions=None |
| ) |
|
|