""" 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) }