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, **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 [] 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 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.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.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) 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.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}) 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}) 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 )