SpliNet / modeling_splinet.py
Angshul's picture
Upload repaired SpliNet 2B-token pretrained model
4ec5e47 verified
Raw History Blame Contribute Delete
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,
)