laya-tr / modeling_laya.py
TurkishCodeMan's picture
Release Turkish Laya-TR non-autoregressive decision model with native AutoModel support
6b96146 verified
Raw
History Blame Contribute Delete
14.2 kB
"""
Laya-TR: Non-Autoregressive Turkish Decision & Reasoning Model
Hugging Face PreTrainedModel uyumlu mimari tanımı.
"""
import math
import time
from typing import Any, Dict, List, Optional, Tuple, Union
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers import PreTrainedModel, AutoTokenizer
try:
from .configuration_laya import LayaConfig
except ImportError:
from configuration_laya import LayaConfig
# -----------------------------------------------------------------------------
# 1. RoPE (Rotary Position Embeddings)
# -----------------------------------------------------------------------------
def rotate_half(x: torch.Tensor) -> torch.Tensor:
x1 = x[..., : x.shape[-1] // 2]
x2 = x[..., x.shape[-1] // 2 :]
return torch.cat((-x2, x1), dim=-1)
def apply_rotary_pos_emb(q: torch.Tensor, k: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
orig_dtype = q.dtype
q_float = q.float()
k_float = k.float()
q_out = (q_float * cos) + (rotate_half(q_float) * sin)
k_out = (k_float * cos) + (rotate_half(k_float) * sin)
return q_out.to(orig_dtype), k_out.to(orig_dtype)
class ModernBertRotaryEmbedding(nn.Module):
def __init__(self, config: LayaConfig):
super().__init__()
self.dim = config.hidden_size // config.num_attention_heads
self.max_seq_len = config.max_position_embeddings
self.theta = config.rope_theta
inv_freq = 1.0 / (self.theta ** (torch.arange(0, self.dim, 2, dtype=torch.float32) / self.dim))
self.register_buffer("inv_freq", inv_freq, persistent=False)
def forward(self, x: torch.Tensor, seq_len: int) -> Tuple[torch.Tensor, torch.Tensor]:
t = torch.arange(seq_len, device=x.device, dtype=torch.float32)
freqs = torch.outer(t, self.inv_freq.to(device=x.device))
emb = torch.cat((freqs, freqs), dim=-1)
cos = emb.cos().unsqueeze(0).unsqueeze(1)
sin = emb.sin().unsqueeze(0).unsqueeze(1)
return cos.to(x.dtype), sin.to(x.dtype)
# -----------------------------------------------------------------------------
# 2. ModernBERT Embeddings & MLP
# -----------------------------------------------------------------------------
class ModernBertEmbeddings(nn.Module):
def __init__(self, config: LayaConfig):
super().__init__()
self.tok_embeddings = nn.Embedding(
config.vocab_size, config.hidden_size, padding_idx=config.pad_token_id
)
self.norm = nn.LayerNorm(config.hidden_size, eps=config.norm_eps, bias=config.norm_bias)
def forward(self, input_ids: torch.Tensor) -> torch.Tensor:
return self.norm(self.tok_embeddings(input_ids))
class ModernBertMLP(nn.Module):
def __init__(self, config: LayaConfig):
super().__init__()
self.Wi = nn.Linear(config.hidden_size, config.intermediate_size * 2, bias=config.mlp_bias)
self.Wo = nn.Linear(config.intermediate_size, config.hidden_size, bias=config.mlp_bias)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
input_gate, hidden = self.Wi(hidden_states).chunk(2, dim=-1)
return self.Wo(F.gelu(input_gate) * hidden)
# -----------------------------------------------------------------------------
# 3. ModernBERT Attention & Encoder Layer
# -----------------------------------------------------------------------------
class ModernBertAttention(nn.Module):
def __init__(self, config: LayaConfig, layer_idx: int):
super().__init__()
self.hidden_size = config.hidden_size
self.num_heads = config.num_attention_heads
self.head_dim = self.hidden_size // self.num_heads
self.layer_idx = layer_idx
self.is_global = (layer_idx % config.global_attn_every_n_layers == 0)
self.local_window = config.local_attention
self.Wqkv = nn.Linear(config.hidden_size, 3 * config.hidden_size, bias=config.attention_bias)
self.Wo = nn.Linear(config.hidden_size, config.hidden_size, bias=config.attention_bias)
def forward(
self,
hidden_states: torch.Tensor,
position_embeddings: Tuple[torch.Tensor, torch.Tensor],
attention_mask: Optional[torch.Tensor] = None
) -> torch.Tensor:
B, S, _ = hidden_states.shape
cos, sin = position_embeddings
qkv = self.Wqkv(hidden_states)
q, k, v = qkv.chunk(3, dim=-1)
q = q.view(B, S, self.num_heads, self.head_dim).transpose(1, 2)
k = k.view(B, S, self.num_heads, self.head_dim).transpose(1, 2)
v = v.view(B, S, self.num_heads, self.head_dim).transpose(1, 2)
q, k = apply_rotary_pos_emb(q, k, cos, sin)
scale = 1.0 / math.sqrt(self.head_dim)
attn_scores = torch.matmul(q, k.transpose(-2, -1)) * scale
if not self.is_global and self.local_window > 0:
row_idx = torch.arange(S, device=hidden_states.device).unsqueeze(1)
col_idx = torch.arange(S, device=hidden_states.device).unsqueeze(0)
sliding_mask = (col_idx < (row_idx - self.local_window)) | (col_idx > (row_idx + self.local_window))
attn_scores = attn_scores.masked_fill(sliding_mask.unsqueeze(0).unsqueeze(0), -1e4)
if attention_mask is not None:
if attention_mask.dim() == 2:
pad_mask = attention_mask.bool().unsqueeze(1).unsqueeze(2)
else:
pad_mask = attention_mask.bool()
attn_scores = attn_scores.masked_fill(~pad_mask, -1e4)
attn_weights = F.softmax(attn_scores, dim=-1, dtype=torch.float32).to(q.dtype)
attn_out = torch.matmul(attn_weights, v)
attn_out = attn_out.transpose(1, 2).contiguous().view(B, S, self.hidden_size)
return self.Wo(attn_out)
class ModernBertEncoderLayer(nn.Module):
def __init__(self, config: LayaConfig, layer_idx: int):
super().__init__()
self.layer_idx = layer_idx
if layer_idx == 0:
self.attn_norm = nn.Identity()
else:
self.attn_norm = nn.LayerNorm(config.hidden_size, eps=config.norm_eps, bias=config.norm_bias)
self.attn = ModernBertAttention(config, layer_idx=layer_idx)
self.mlp_norm = nn.LayerNorm(config.hidden_size, eps=config.norm_eps, bias=config.norm_bias)
self.mlp = ModernBertMLP(config)
def forward(
self,
hidden_states: torch.Tensor,
position_embeddings: Tuple[torch.Tensor, torch.Tensor],
attention_mask: Optional[torch.Tensor] = None
) -> torch.Tensor:
attn_out = self.attn(
self.attn_norm(hidden_states),
position_embeddings=position_embeddings,
attention_mask=attention_mask
)
hidden_states = hidden_states + attn_out
mlp_out = self.mlp(self.mlp_norm(hidden_states))
hidden_states = hidden_states + mlp_out
return hidden_states
class ModernBertEncoder(nn.Module):
def __init__(self, config: LayaConfig):
super().__init__()
self.config = config
self.embeddings = ModernBertEmbeddings(config)
self.rotary_emb = ModernBertRotaryEmbedding(config)
self.layers = nn.ModuleList([
ModernBertEncoderLayer(config, layer_idx=l)
for l in range(config.num_hidden_layers)
])
self.final_norm = nn.LayerNorm(config.hidden_size, eps=config.norm_eps, bias=config.norm_bias)
def forward(self, input_ids: torch.Tensor, attention_mask: Optional[torch.Tensor] = None) -> torch.Tensor:
B, S = input_ids.shape
hidden_states = self.embeddings(input_ids)
position_embeddings = self.rotary_emb(hidden_states, seq_len=S)
for layer in self.layers:
hidden_states = layer(
hidden_states,
position_embeddings=position_embeddings,
attention_mask=attention_mask
)
hidden_states = self.final_norm(hidden_states)
return hidden_states
# -----------------------------------------------------------------------------
# 4. Decision Transformer Head, Scorer & Act Head
# -----------------------------------------------------------------------------
class DecisionTransformerHead(nn.Module):
def __init__(self, config: LayaConfig):
super().__init__()
d = config.hidden_size
nhead = config.num_attention_heads
d_ff = config.head_ff_dim
self.layers = nn.ModuleList([
nn.TransformerEncoderLayer(
d_model=d,
nhead=nhead,
dim_feedforward=d_ff,
dropout=0.1,
batch_first=True,
norm_first=True
)
for _ in range(config.head_layers)
])
def forward(self, x: torch.Tensor, attention_mask: Optional[torch.Tensor] = None) -> torch.Tensor:
pad_mask = ~attention_mask.bool() if attention_mask is not None else None
for layer in self.layers:
x = layer(x, src_key_padding_mask=pad_mask)
return x
# -----------------------------------------------------------------------------
# 5. Hugging Face PreTrainedModel Uyumlu LayaDecisionModel
# -----------------------------------------------------------------------------
class LayaDecisionModel(PreTrainedModel):
config_class = LayaConfig
base_model_prefix = "laya"
supports_gradient_checkpointing = True
def __init__(self, config: LayaConfig):
super().__init__(config)
d = config.hidden_size
self.encoder = ModernBertEncoder(config)
self.type_emb = nn.Embedding(config.num_question_types, d)
self.head = DecisionTransformerHead(config) if config.head_layers > 0 else None
self.scorer = nn.Sequential(
nn.LayerNorm(d),
nn.Linear(d, d),
nn.GELU(),
nn.Linear(d, 1)
)
self.act_head = nn.Sequential(
nn.Linear(d + 4, 256),
nn.GELU(),
nn.Linear(256, config.n_act)
)
self.register_buffer("temperature", torch.ones(3))
self.post_init()
def forward(
self,
input_ids: torch.Tensor,
attention_mask: torch.Tensor,
marker_pos: torch.Tensor,
marker_mask: torch.Tensor,
qtype: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor]:
h = self.encoder(input_ids=input_ids, attention_mask=attention_mask)
h = h + self.type_emb(qtype)[:, None, :]
if self.head is not None:
h = self.head(h, attention_mask=attention_mask)
idx = marker_pos.clamp(min=0)[:, :, None].expand(-1, -1, h.size(-1))
m = torch.gather(h, 1, idx)
logits = self.scorer(m).squeeze(-1).float()
logits = logits.masked_fill(~marker_mask, -1e4)
p = torch.softmax(logits.detach(), dim=-1)
k = marker_mask.sum(-1).clamp(min=2).float()
ent = -(p * torch.log(p.clamp_min(1e-9))).sum(-1) / torch.log(k)
if p.size(-1) >= 2:
top2 = p.topk(2, dim=-1).values
else:
top1 = p.topk(1, dim=-1).values
top2 = torch.cat([top1, torch.zeros_like(top1)], dim=-1)
feats = torch.stack([top2[:, 0], top2[:, 0] - top2[:, 1], ent, k / 255.0], dim=-1)
pooled = h[:, 0]
act_input = torch.cat([pooled, feats.to(dtype=h.dtype)], dim=-1)
act_logits = self.act_head(act_input)
return logits, act_logits
@torch.no_grad()
def decide(
self,
question: str,
options: Union[List[str], Dict[str, str]],
tokenizer: Optional[AutoTokenizer] = None,
context: Optional[str] = None
) -> Dict[str, Any]:
"""
Kullanıcıların tek satırda AutoModel üzerinden sub-10ms karar almasını sağlar.
"""
if tokenizer is None:
tokenizer = AutoTokenizer.from_pretrained("jhu-clsp/mmBERT-base")
t0 = time.perf_counter()
device = next(self.parameters()).device
if isinstance(options, dict):
opt_labels = list(options.keys())
opt_texts = [f"{k}: {v}" if v else k for k, v in options.items()]
else:
opt_labels = [chr(65 + i) for i in range(len(options))]
opt_texts = [f"{lbl}: {opt}" for lbl, opt in zip(opt_labels, options)]
mask_tok = tokenizer.mask_token
head_ids = tokenizer(f"choice question: {question}", add_special_tokens=False)["input_ids"]
opt_ids = []
for text in opt_texts:
opt_ids.append(tokenizer(f"{mask_tok} {text}", add_special_tokens=False)["input_ids"])
cls_id = tokenizer.cls_token_id or 1
sep_id = tokenizer.sep_token_id or 1
seq = [cls_id] + head_ids + [sep_id]
markers = []
for o_ids in opt_ids:
markers.append(len(seq))
seq.extend(o_ids)
seq.append(sep_id)
if context:
ctx_ids = tokenizer(str(context), add_special_tokens=False)["input_ids"][:512]
seq.extend(ctx_ids)
seq.append(sep_id)
input_ids = torch.tensor([seq], dtype=torch.long, device=device)
attention_mask = torch.ones_like(input_ids)
marker_pos = torch.tensor([markers], dtype=torch.long, device=device)
marker_mask = torch.ones_like(marker_pos, dtype=torch.bool)
qtype = torch.tensor([0], dtype=torch.long, device=device)
logits, act_logits = self(
input_ids=input_ids,
attention_mask=attention_mask,
marker_pos=marker_pos,
marker_mask=marker_mask,
qtype=qtype
)
probs = F.softmax(logits[0], dim=-1).cpu().tolist()
best_idx = int(torch.argmax(logits[0]).item())
elapsed_ms = (time.perf_counter() - t0) * 1000
prob_map = {lbl: round(p, 4) for lbl, p in zip(opt_labels, probs)}
return {
"prediction": opt_labels[best_idx],
"selected_option": opt_texts[best_idx],
"confidence": round(probs[best_idx], 4),
"probabilities": prob_map,
"latency_ms": round(elapsed_ms, 2)
}