Instructions to use Angshul/SpliNet with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use Angshul/SpliNet with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("fill-mask", model="Angshul/SpliNet", trust_remote_code=True)# pip install -U transformers accelerate # Load model directly from transformers import AutoModelForMaskedLM model = AutoModelForMaskedLM.from_pretrained("Angshul/SpliNet", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download modeling_splinet.py from Angshul/SpliNet: direct link, hf CLI and curl.
- Browser
- Download file 8.01 kB
-
https://huggingface.co/Angshul/SpliNet/resolve/main/modeling_splinet.py
- Command line
-
hf download hf://Angshul/SpliNet/modeling_splinet.py
-
curl -L -o modeling_splinet.py https://huggingface.co/Angshul/SpliNet/resolve/main/modeling_splinet.py
8.01 kB
| 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, | |
| ) | |