Release Turkish Laya-TR non-autoregressive decision model with native AutoModel support
6b96146 verified | """ | |
| 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 | |
| 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) | |
| } | |