factorized-null / tolm.py
miguelcsx's picture
main chck_88M
7c8a736 verified
Raw
History Blame Contribute Delete
35.9 kB
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
)