bert-base-ner-mlx / bert_ner_mlx.py
masahiroid's picture
Upload folder using huggingface_hub
b28b22c verified
Raw History Blame Contribute Delete
3.47 kB
"""MLX reimplementation of dslim/bert-base-NER (standard BERT-base encoder +
a per-token linear classifier head for named-entity recognition).
Standard post-norm BERT, not covered by `mlx-embeddings` (which targets
embedding/reranker pooling outputs, not token classification), so
reimplemented directly against the HF `transformers` BERT source.
"""
import math
import mlx.core as mx
import mlx.nn as nn
HIDDEN_SIZE = 768
NUM_HEADS = 12
HEAD_DIM = HIDDEN_SIZE // NUM_HEADS
INTERMEDIATE_SIZE = 3072
NUM_LAYERS = 12
NUM_LABELS = 9
VOCAB_SIZE = 28996
MAX_POSITION_EMBEDDINGS = 512
TYPE_VOCAB_SIZE = 2
LN_EPS = 1e-12
class BertLayer(nn.Module):
def __init__(self):
super().__init__()
self.q = nn.Linear(HIDDEN_SIZE, HIDDEN_SIZE)
self.k = nn.Linear(HIDDEN_SIZE, HIDDEN_SIZE)
self.v = nn.Linear(HIDDEN_SIZE, HIDDEN_SIZE)
self.attn_out = nn.Linear(HIDDEN_SIZE, HIDDEN_SIZE)
self.attn_layernorm = nn.LayerNorm(HIDDEN_SIZE, eps=LN_EPS)
self.intermediate = nn.Linear(HIDDEN_SIZE, INTERMEDIATE_SIZE)
self.output = nn.Linear(INTERMEDIATE_SIZE, HIDDEN_SIZE)
self.output_layernorm = nn.LayerNorm(HIDDEN_SIZE, eps=LN_EPS)
def __call__(self, x, attn_bias=None):
b, n, c = x.shape
q = self.q(x).reshape(b, n, NUM_HEADS, HEAD_DIM).transpose(0, 2, 1, 3)
k = self.k(x).reshape(b, n, NUM_HEADS, HEAD_DIM).transpose(0, 2, 1, 3)
v = self.v(x).reshape(b, n, NUM_HEADS, HEAD_DIM).transpose(0, 2, 1, 3)
attn = mx.fast.scaled_dot_product_attention(
q, k, v, scale=1.0 / math.sqrt(HEAD_DIM), mask=attn_bias
)
attn = attn.transpose(0, 2, 1, 3).reshape(b, n, c)
attn = self.attn_out(attn)
x = self.attn_layernorm(x + attn)
h = self.output(nn.gelu(self.intermediate(x)))
x = self.output_layernorm(x + h)
return x
class BertNerMLX(nn.Module):
def __init__(self):
super().__init__()
self.word_embeddings = mx.zeros((VOCAB_SIZE, HIDDEN_SIZE))
self.position_embeddings = mx.zeros((MAX_POSITION_EMBEDDINGS, HIDDEN_SIZE))
self.token_type_embeddings = mx.zeros((TYPE_VOCAB_SIZE, HIDDEN_SIZE))
self.embeddings_layernorm = nn.LayerNorm(HIDDEN_SIZE, eps=LN_EPS)
self.layers = [BertLayer() for _ in range(NUM_LAYERS)]
self.classifier = nn.Linear(HIDDEN_SIZE, NUM_LABELS)
def __call__(self, input_ids, attention_mask=None, token_type_ids=None):
b, n = input_ids.shape
word_emb = self.word_embeddings[input_ids]
pos_ids = mx.arange(n)
pos_emb = self.position_embeddings[pos_ids]
if token_type_ids is None:
token_type_ids = mx.zeros((b, n), dtype=mx.int32)
type_emb = self.token_type_embeddings[token_type_ids]
x = word_emb + pos_emb[None] + type_emb
x = self.embeddings_layernorm(x)
attn_bias = None
if attention_mask is not None:
# Use a large-but-finite negative value that fits in fp16 (max ~65504).
# A literal like -1e9 overflows fp16 to -inf, and 0 * -inf = NaN for
# the unmasked (mask=1) positions where this term should be exactly 0.
neg = mx.array(-1e4, dtype=x.dtype)
attn_bias = (1 - attention_mask[:, None, None, :].astype(x.dtype)) * neg
for layer in self.layers:
x = layer(x, attn_bias=attn_bias)
logits = self.classifier(x)
return logits