LFM2.5-Encoder-350M-PII-Detector / modeling_phase2_tc.py
iamleonie's picture
Initial public commit
b8c9cf3
Raw
History Blame Contribute Delete
6.55 kB
"""Token-classification head on the PHASE-2 bidirectional LFM2.5 MLM encoder
(LiquidAI/mlm_phase2_bidir2_step140800, "theirs" bidir variant: full gated shortconv
made symmetric + non-causal attention). Self-contained for trust_remote_code.
The encoder section is inlined VERBATIM from the backbone's shipped
`modeling_lfm2_bidir_theirs.py` (so the trained conv in_proj/out_proj/conv weights are
used exactly as trained — unlike our v13 vendored patch which dropped in_proj/out_proj).
On top sits our BIOES classifier head + class-weighted label-smoothed CE.
"""
from __future__ import annotations
from typing import Optional
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers.modeling_outputs import TokenClassifierOutput
from transformers.models.lfm2 import modeling_lfm2 as _lfm2_mod
from transformers.models.lfm2.configuration_lfm2 import Lfm2Config
from transformers.models.lfm2.modeling_lfm2 import (
Lfm2Attention, Lfm2Model, Lfm2PreTrainedModel, Lfm2ShortConv, apply_mask_to_padding_states,
)
# ---- bidirectional patches (verbatim from the backbone's modeling_lfm2_bidir_theirs) ----
def _bidirectional_mask(config, input_embeds=None, attention_mask=None, cache_position=None,
past_key_values=None, position_ids=None, **kwargs):
if input_embeds is None:
input_embeds = kwargs.get("inputs_embeds")
if config._attn_implementation == "flash_attention_2":
if attention_mask is not None and not attention_mask.all():
return attention_mask
return None
device = input_embeds.device; dtype = input_embeds.dtype
bsz, q_len = input_embeds.shape[:2]
past = past_key_values.get_seq_length() if past_key_values is not None else 0
kv_len = past + q_len
mask = torch.zeros((bsz, 1, q_len, kv_len), device=device, dtype=dtype)
if attention_mask is not None:
cur_len = attention_mask.size(-1)
key_pad_flags = (attention_mask == 0).to(device=device, dtype=torch.float32)
pad_vec = torch.zeros((bsz, kv_len), device=device, dtype=torch.float32)
if cur_len > 0:
pad_vec[:, past:past + cur_len] = key_pad_flags * -1e9
mask = mask + pad_vec.to(dtype)[:, None, None, :]
return mask
def _noncausal_shortconv_forward(self, hidden_states, past_key_values=None, cache_position=None,
attention_mask=None, **kwargs):
x = apply_mask_to_padding_states(hidden_states, attention_mask)
BCx = self.in_proj(x).transpose(-1, -2)
B, C, x = BCx.chunk(3, dim=-2)
Bx = B * x
k = self.conv.weight.shape[-1]; pad = k // 2
conv_out = F.conv1d(Bx, weight=self.conv.weight, bias=self.conv.bias,
stride=1, padding=pad, dilation=1, groups=Bx.shape[1])
if conv_out.shape[-1] > Bx.shape[-1]:
conv_out = conv_out[..., :Bx.shape[-1]]
elif conv_out.shape[-1] < Bx.shape[-1]:
conv_out = F.pad(conv_out, (0, Bx.shape[-1] - conv_out.shape[-1]))
y = C * conv_out
y = y.transpose(-1, -2).contiguous()
return self.out_proj(y)
def _shortconv_forward(self, *args, **kwargs):
return self.slow_forward(*args, **kwargs)
_PATCHED = False
def _install_patches():
global _PATCHED
if _PATCHED:
return
_lfm2_mod.create_causal_mask = _bidirectional_mask
Lfm2ShortConv.slow_forward = _noncausal_shortconv_forward
Lfm2ShortConv.forward = _shortconv_forward
_PATCHED = True
_install_patches()
def _set_attention_noncausal(model):
for m in model.modules():
if isinstance(m, Lfm2Attention):
m.is_causal = False
class Lfm2BidirectionalModel_theirs(Lfm2Model):
def __init__(self, config):
_install_patches()
super().__init__(config)
_set_attention_noncausal(self)
# ---- token-classification head (BIOES) ----
class Lfm2BidirP2ForTokenClassification(Lfm2PreTrainedModel):
config_class = Lfm2Config
base_model_prefix = "lfm2"
def __init__(self, config: Lfm2Config):
_install_patches()
config = type(config).from_dict({**config.to_dict(), "use_cache": False})
super().__init__(config)
self.num_labels = config.num_labels
self.lfm2 = Lfm2BidirectionalModel_theirs(config)
cd = getattr(config, "classifier_dropout", None)
self.dropout = nn.Dropout(cd if cd is not None else 0.1)
self.classifier = nn.Linear(config.hidden_size, config.num_labels)
self.register_buffer("class_weights", torch.ones(config.num_labels), persistent=True)
self.label_smoothing = float(getattr(config, "label_smoothing", 0.0))
self.post_init()
def set_class_weights(self, weights: torch.Tensor) -> None:
if weights.shape != (self.num_labels,):
raise ValueError(f"class_weights must be ({self.num_labels},), got {tuple(weights.shape)}")
self.class_weights = weights.to(device=self.class_weights.device, dtype=torch.float32)
def get_input_embeddings(self):
return self.lfm2.embed_tokens
def set_input_embeddings(self, value):
self.lfm2.embed_tokens = value
def forward(self, input_ids=None, attention_mask=None, position_ids=None,
inputs_embeds=None, labels=None, output_hidden_states=None,
output_attentions=None, return_dict=None, **kwargs) -> TokenClassifierOutput:
return_dict = True if return_dict is None else return_dict
outputs = self.lfm2(input_ids=input_ids, attention_mask=attention_mask,
position_ids=position_ids, inputs_embeds=inputs_embeds,
use_cache=False, output_attentions=output_attentions,
output_hidden_states=output_hidden_states, return_dict=True)
hidden = self.dropout(outputs.last_hidden_state)
logits = self.classifier(hidden)
loss = None
if labels is not None:
cw = self.class_weights.float()
if not torch.isfinite(cw).all():
cw = torch.ones_like(cw)
loss = F.cross_entropy(logits.view(-1, self.num_labels).float(), labels.view(-1),
weight=cw, ignore_index=-100, label_smoothing=self.label_smoothing)
if not return_dict:
out = (logits,) + outputs[1:]
return ((loss,) + out) if loss is not None else out
return TokenClassifierOutput(loss=loss, logits=logits,
hidden_states=outputs.hidden_states, attentions=outputs.attentions)