import math import torch import torch.nn as nn import torch.nn.functional as F from transformers.modeling_outputs import ( BaseModelOutput, MaskedLMOutput, ) from transformers.models.fnet.modeling_fnet import ( FNetEmbeddings, FNetIntermediate, FNetOnlyMLMHead, FNetOutput, FNetPreTrainedModel, ) from .configuration_splinet import SpliNetConfig class SplineMixer(nn.Module): def __init__(self, config): super().__init__() self.radius = int( config.splinet_radius ) z = -3.0 + 2.0 * math.sqrt(2.0) # Fixed analytic spline coefficients. # # IMPORTANT: # Do NOT register these as a non-persistent tensor buffer. # Hugging Face low-memory/meta-device loading can materialize # such a buffer without its analytically initialized values. # # Store the 33 coefficients as ordinary Python floats instead. # They are recreated on the actual input device and dtype in # forward(). They are not learned model state. self.kernel_values = tuple( float( math.sqrt(2.0) * (z ** abs(k)) ) for k in range( -self.radius, self.radius + 1, ) ) def forward(self, x): if x.shape[1] != 512: raise ValueError( "This released SpliNet checkpoint was pretrained " "and validated on fixed 512-token blocks. " f"Received sequence length {x.shape[1]}. " "Tokenize/pack the input to exactly 512 tokens." ) d = x.shape[-1] y = ( x.transpose(1, 2) .contiguous() ) y = F.pad( y, ( self.radius, self.radius, ), mode="reflect", ) kernel = ( torch.tensor( self.kernel_values, device=y.device, dtype=y.dtype, ) .view(1, 1, -1) .expand(d, 1, -1) .contiguous() ) y = F.conv1d( y, kernel, groups=d, ) return ( y.transpose(1, 2) .contiguous() ) class SpliNetMixingBlock(nn.Module): def __init__(self, config): super().__init__() self.mixer = SplineMixer( config ) self.LayerNorm = nn.LayerNorm( config.hidden_size, eps=config.layer_norm_eps, ) def forward(self, x): return self.LayerNorm( x + self.mixer(x) ) class SpliNetLayer(nn.Module): def __init__(self, config): super().__init__() self.mixing = ( SpliNetMixingBlock( config ) ) self.intermediate = ( FNetIntermediate( config ) ) self.output = ( FNetOutput( config ) ) def forward(self, x): x = self.mixing(x) return self.output( self.intermediate(x), x, ) class SpliNetEncoder(nn.Module): def __init__(self, config): super().__init__() self.layer = nn.ModuleList( [ SpliNetLayer(config) for _ in range( config.num_hidden_layers ) ] ) def forward( self, x, output_hidden_states=False, ): hidden_states = ( () if output_hidden_states else None ) for layer in self.layer: if output_hidden_states: hidden_states += (x,) x = layer(x) if output_hidden_states: hidden_states += (x,) return BaseModelOutput( last_hidden_state=x, hidden_states=hidden_states, ) class SpliNetModel(FNetPreTrainedModel): config_class = SpliNetConfig base_model_prefix = "splinet" def __init__(self, config): super().__init__(config) self.embeddings = ( FNetEmbeddings(config) ) self.encoder = ( SpliNetEncoder(config) ) self.post_init() def get_input_embeddings(self): return ( self.embeddings .word_embeddings ) def set_input_embeddings( self, value, ): self.embeddings.word_embeddings = value def forward( self, input_ids=None, token_type_ids=None, position_ids=None, inputs_embeds=None, output_hidden_states=False, **kwargs, ): if input_ids is not None: shape = input_ids.shape device = input_ids.device elif inputs_embeds is not None: shape = inputs_embeds.shape[:-1] device = inputs_embeds.device else: raise ValueError( "input_ids or inputs_embeds required." ) if shape[1] != 512: raise ValueError( "This SpliNet checkpoint requires exactly " f"512 tokens; received {shape[1]}." ) if token_type_ids is None: token_type_ids = torch.zeros( shape, dtype=torch.long, device=device, ) x = self.embeddings( input_ids=input_ids, token_type_ids=token_type_ids, position_ids=position_ids, inputs_embeds=inputs_embeds, ) return self.encoder( x, output_hidden_states=output_hidden_states, ) class SpliNetForMaskedLM(FNetPreTrainedModel): config_class = SpliNetConfig base_model_prefix = "splinet" _tied_weights_keys = { "cls.predictions.decoder.bias": "cls.predictions.bias", "cls.predictions.decoder.weight": "splinet.embeddings.word_embeddings.weight", } def __init__(self, config): super().__init__(config) self.splinet = ( SpliNetModel(config) ) self.cls = ( FNetOnlyMLMHead(config) ) self.post_init() if config.tie_word_embeddings: ( self.cls .predictions .decoder.weight ) = ( self.splinet .embeddings .word_embeddings .weight ) def get_input_embeddings(self): return ( self.splinet .embeddings .word_embeddings ) def set_input_embeddings( self, value, ): ( self.splinet .embeddings .word_embeddings ) = value def get_output_embeddings(self): return ( self.cls .predictions .decoder ) def set_output_embeddings( self, value, ): self.cls.predictions.decoder = value def forward( self, input_ids=None, labels=None, **kwargs, ): out = self.splinet( input_ids=input_ids, **kwargs, ) logits = self.cls( out.last_hidden_state ) loss = None if labels is not None: loss = F.cross_entropy( logits.reshape( -1, self.config.vocab_size, ), labels.reshape(-1), ignore_index=-100, ) return MaskedLMOutput( loss=loss, logits=logits, hidden_states=out.hidden_states, )