Instructions to use miguelcsx/prism-control with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use miguelcsx/prism-control with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("fill-mask", model="miguelcsx/prism-control", trust_remote_code=True)# Load model directly from transformers import AutoModelForMaskedLM model = AutoModelForMaskedLM.from_pretrained("miguelcsx/prism-control", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
| 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 | |
| ) | |
| 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 | |
| ) | |