"""Configuration for DZAIR encoder family (base + small).""" from __future__ import annotations from typing import Any from transformers import PretrainedConfig _BASE_HIDDEN_SIZE = 768 _BASE_LAYERS = 12 _SMALL_HIDDEN_SIZE = 384 _SMALL_LAYERS = 6 class DzairConfig(PretrainedConfig): """DZAIR encoder configuration. Every architectural choice is a declared field so ``config.json`` round-trips exactly (base and small share this class, never a hidden ``arch`` object). ``**kwargs`` forwards only transformers-managed keys (e.g. ``transformers_version``) to ``PretrainedConfig``. Field names match the published `config.json`. Two sizes share this config: - base: 12Lx768 (discriminator, grouped-query 12Q/4KV) + 3Lx384 generator, shared embeddings - small: 6Lx384 (discriminator, grouped-query 6Q/2KV) + 3Lx384 generator, shared embeddings """ model_type = "dzair" def __init__( self, vocab_size: int = 48000, hidden_size: int = 768, intermediate_size: int = 1792, num_attention_heads: int = 12, num_key_value_heads: int = 0, num_hidden_layers: int = 12, num_generator_layers: int = 3, generator_hidden_size: int = 384, generator_intermediate_size: int = 1024, max_position_embeddings: int = 512, rope_theta: float = 10000.0, hidden_dropout_prob: float = 0.1, attention_probs_dropout_prob: float = 0.1, layer_norm_eps: float = 1e-5, pad_token_id: int = 0, cls_token_id: int = 2, sep_token_id: int = 3, mask_token_id: int = 4, tie_word_embeddings: bool = True, share_generator_embeddings: bool = False, qk_norm: bool = False, **kwargs: Any, ) -> None: if hidden_size % num_attention_heads != 0: msg = f"hidden_size {hidden_size} must split over {num_attention_heads} heads" raise ValueError(msg) if num_key_value_heads == 0: num_key_value_heads = num_attention_heads if num_attention_heads % num_key_value_heads != 0: msg = ( f"{num_attention_heads} query heads must split over " f"{num_key_value_heads} key-value heads" ) raise ValueError(msg) if generator_hidden_size % 64 != 0: msg = f"generator_hidden_size {generator_hidden_size} must be a multiple of 64" raise ValueError(msg) self.vocab_size = vocab_size self.hidden_size = hidden_size self.intermediate_size = intermediate_size self.num_attention_heads = num_attention_heads self.num_key_value_heads = num_key_value_heads self.num_hidden_layers = num_hidden_layers self.num_generator_layers = num_generator_layers self.generator_hidden_size = generator_hidden_size self.generator_intermediate_size = generator_intermediate_size self.max_position_embeddings = max_position_embeddings self.rope_theta = rope_theta self.hidden_dropout_prob = hidden_dropout_prob self.attention_probs_dropout_prob = attention_probs_dropout_prob self.layer_norm_eps = layer_norm_eps self.share_generator_embeddings = share_generator_embeddings self.qk_norm = qk_norm super().__init__( pad_token_id=pad_token_id, cls_token_id=cls_token_id, sep_token_id=sep_token_id, tie_word_embeddings=tie_word_embeddings, **kwargs, ) # Ensure mask_token_id and explicit IDs are preserved as ints self.pad_token_id = pad_token_id self.cls_token_id = cls_token_id self.sep_token_id = sep_token_id self.mask_token_id = mask_token_id @property def head_size(self) -> int: return self.hidden_size // self.num_attention_heads @property def generator_num_heads(self) -> int: """Generator query heads at head_dim 64 (always divides, checked above).""" return self.generator_hidden_size // 64 @property def kv_dim(self) -> int: """Key/value width: key-value heads at the trunk head_dim.""" return self.num_key_value_heads * self.head_size @property def is_base(self) -> bool: return self.hidden_size == _BASE_HIDDEN_SIZE and self.num_hidden_layers == _BASE_LAYERS @property def is_small(self) -> bool: return self.hidden_size == _SMALL_HIDDEN_SIZE and self.num_hidden_layers == _SMALL_LAYERS # Predefined configurations. Both sizes share the generator embedding table # with the discriminator (GDES, DeBERTaV3) and apply QK-norm; the FFN # intermediate is 128-aligned for tensor cores (1792 = 14x128, 1024 = 8x128). DZAIR_BASE_CONFIG = DzairConfig( num_key_value_heads=4, share_generator_embeddings=True, qk_norm=True, ) DZAIR_SMALL_CONFIG = DzairConfig( hidden_size=384, intermediate_size=1024, num_attention_heads=6, num_key_value_heads=2, num_hidden_layers=6, num_generator_layers=3, generator_hidden_size=384, generator_intermediate_size=1024, share_generator_embeddings=True, qk_norm=True, ) __all__ = ["DZAIR_BASE_CONFIG", "DZAIR_SMALL_CONFIG", "DzairConfig"] """DZAIR encoder: RTD + GDES (DeBERTaV3 objective) with ModernBERT-speed architecture. Architecture: pre-RMSNorm, RoPE, SwiGLU, fused scaled-dot-product attention (FlashAttention-2 path when available) attending globally, single-chunk sequences. Objective: RTD on all tokens. Generator MLM corrupts; GDES detaches generator embeddings. """ import copy import hashlib import math from dataclasses import dataclass from pathlib import Path from typing import Any, ClassVar import torch from torch import Tensor, _dynamo, nn from torch.nn import functional from torch.utils import checkpoint as checkpoint_utils from transformers import PretrainedConfig, PreTrainedModel from transformers.utils.generic import ModelOutput IGNORE_INDEX = -100 # ELECTRA (Clark et al., 2020, §3.3): small models weight the discriminator # loss at 50 relative to the generator MLM loss. RTD_LOSS_WEIGHT = 50.0 # BERT 80/10/10 corruption splits (Devlin et al., 2019): below REPLACE the # token becomes [MASK], below REPLACE+RANDOM it becomes a random vocab id, # otherwise it is kept (but still predicted by the generator). MASK_REPLACE_CUTOFF = 0.8 MASK_RANDOM_CUTOFF = 0.9 def _apply_rope(x: Tensor, cos: Tensor, sin: Tensor) -> Tensor: """Apply rotary positional embeddings to half the head dim.""" x1, x2 = x.chunk(2, dim=-1) return torch.cat((x1 * cos - x2 * sin, x1 * sin + x2 * cos), dim=-1) def _build_rope_cache( max_seq_len: int, head_dim: int, theta: float, device: torch.device ) -> tuple[Tensor, Tensor]: """Build RoPE cos/sin cache for sequence length up to max_seq_len.""" inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2, device=device).float() / head_dim)) t = torch.arange(max_seq_len, device=device).float() freqs = torch.outer(t, inv_freq) cos = freqs.cos().to(torch.get_default_dtype()) sin = freqs.sin().to(torch.get_default_dtype()) return cos, sin class RMSNorm(nn.Module): """Root Mean Square Layer Normalization (affine weight, no bias). The statistic is computed in float32 and the result cast back. In half precision the square overflows: trunk activations reach 316, and 316 squared is 99,856 against a float16 maximum of 65,504, so the mean becomes inf, its reciprocal square root becomes 0, and the encoder returns an all-zero hidden state (measured 2026-09-19 on the released fp16 build). float32 inputs are unaffected — the upcast is a no-op and outputs stay bit-identical. """ def __init__(self, dim: int, eps: float = 1e-5) -> None: super().__init__() self.eps = eps self.weight = nn.Parameter(torch.ones(dim)) def forward(self, x: Tensor) -> Tensor: working = x.float() norm = working.pow(2).mean(dim=-1, keepdim=True).add(self.eps).rsqrt() return (working * norm * self.weight.float()).to(x.dtype) class SwiGLU(nn.Module): """Swish-Gated Linear Unit.""" def forward(self, x: Tensor) -> Tensor: x, gate = x.chunk(2, dim=-1) return x * functional.silu(gate) class FeedForward(nn.Module): """Pre-RMSNorm SwiGLU FFN with dropout.""" def __init__(self, config: DzairConfig) -> None: super().__init__() self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_eps) self.up = nn.Linear(config.hidden_size, 2 * config.intermediate_size, bias=False) self.act = SwiGLU() self.down = nn.Linear(config.intermediate_size, config.hidden_size, bias=False) self.dropout = nn.Dropout(config.hidden_dropout_prob) def forward(self, x: Tensor) -> Tensor: residual = x x = self.norm(x) x = self.up(x) x = self.act(x) x = self.dropout(self.down(x)) return residual + x class Attention(nn.Module): """Grouped-query attention with RoPE: every layer attends globally. Query heads share fewer key/value heads (``num_key_value_heads`` groups). Local-window alternation was cut 2026-09-14: at 512 tokens it saves ~4% wall-clock (measured FLOP arithmetic) while full attention is the literature default every baseline trains — the deviation bought complexity without evidence. Fused projections, bias-free, pre-RMSNorm. """ # Declared so the registered buffers carry a type; register_buffer alone # leaves them untyped for the checker. _cos: Tensor _sin: Tensor def __init__(self, config: DzairConfig) -> None: super().__init__() self.config = config self.num_heads = config.num_attention_heads self.num_kv_heads = config.num_key_value_heads self.head_size = config.head_size self.scale = 1.0 / math.sqrt(self.head_size) # Separate Q and fused KV projections (bias-free for FA-2 compatibility). # KV groups repeat to the query count at forward time. self.q_proj = nn.Linear(config.hidden_size, config.hidden_size, bias=False) self.kv_proj = nn.Linear(config.hidden_size, 2 * config.kv_dim, bias=False) self.out_proj = nn.Linear(config.hidden_size, config.hidden_size, bias=False) self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_eps) # QK-norm (Gemma 2/3 practice): one shared RMSNorm over head_dim # applied to queries and keys before RoPE. Norm-then-rotate is a # fixed convention, not a commutation (rotation mixes dims, so an # affine weight does not commute with it) — the trained weights # bake in this order, so it must never change under them. self.qk_norm: RMSNorm | None = ( RMSNorm(self.head_size, eps=config.layer_norm_eps) if config.qk_norm else None ) self.dropout = nn.Dropout(config.attention_probs_dropout_prob) # RoPE cache: non-persistent buffers, so they stay out of the state # dict but are visible to the ONNX exporter, which warns about plain # attributes assigned during a traced forward. self.register_buffer("_cos", torch.empty(0), persistent=False) self.register_buffer("_sin", torch.empty(0), persistent=False) def _get_rope(self, seq_len: int, device: torch.device) -> tuple[Tensor, Tensor]: if self._cos.numel() == 0 or self._cos.size(0) < seq_len or self._cos.device != device: self._cos, self._sin = _build_rope_cache( max(seq_len, self.config.max_position_embeddings), self.head_size, self.config.rope_theta, device, ) return self._cos[:seq_len], self._sin[:seq_len] def forward( self, x: Tensor, attention_mask: Tensor | None = None, is_causal: bool = False, ) -> Tensor: """x: [B, T, D], attention_mask: [B, T] (1=keep, 0=pad). Returns [B, T, D].""" batch_size, seq_len, _ = x.shape # Pre-norm x_norm = self.norm(x) # Grouped-query projections. q = self.q_proj(x_norm) # [B, T, D] kv = self.kv_proj(x_norm) # [B, T, 2 * kv_dim] k, v = kv.chunk(2, dim=-1) # Reshape for attention: Q [B, H, T, head_dim], K/V [B, KV, T, head_dim]. q = q.view(batch_size, seq_len, self.num_heads, self.head_size).transpose(1, 2) k = k.view(batch_size, seq_len, self.num_kv_heads, self.head_size).transpose(1, 2) v = v.view(batch_size, seq_len, self.num_kv_heads, self.head_size).transpose(1, 2) # Repeat KV groups to the query count (exact: heads split evenly, checked). repeat = self.num_heads // self.num_kv_heads if repeat > 1: k = k.repeat_interleave(repeat, dim=1) v = v.repeat_interleave(repeat, dim=1) if self.qk_norm is not None: q = self.qk_norm(q) k = self.qk_norm(k) # RoPE (cast to the working dtype: an fp32 cache multiplied into bf16 # queries upcasts them and drops out of the fused-attention fast path) cos, sin = self._get_rope(seq_len, x.device) cos = cos.unsqueeze(0).unsqueeze(0).to(x.dtype) # [1, 1, T, head_dim/2] sin = sin.unsqueeze(0).unsqueeze(0).to(x.dtype) q = _apply_rope(q, cos, sin) k = _apply_rope(k, cos, sin) # Scaled dot-product attention. A bool mask (True = attend) keeps the # fused fast path; the old additive float mask did not. attn_mask: Tensor | None = None if attention_mask is not None: attn_mask = attention_mask.to(torch.bool).view(batch_size, 1, 1, seq_len) # Guard against all-False mask rows (all-pad inputs): SDPA under # CUDA/Inductor produces NaNs when a row has zero attendable keys. positions = torch.arange(seq_len, device=x.device) has_key = attn_mask.any(dim=-1, keepdim=True) attn_mask = attn_mask | (~has_key & (positions == 0).view(1, 1, 1, seq_len)) # Use PyTorch's scaled_dot_product_attention (uses FA-2 when available) attn_out = functional.scaled_dot_product_attention( q, k, v, attn_mask=attn_mask, dropout_p=self.config.attention_probs_dropout_prob if self.training else 0.0, is_causal=is_causal, scale=self.scale, ) # Merge heads: [B, H, T, head_dim] -> [B, T, D] attn_out = attn_out.transpose(1, 2).contiguous().view(batch_size, seq_len, -1) # Zero out padding positions so unused positions never dominate # downstream means (the residual still carries the pad embedding; # losses and CLS pooling ignore pads by mask, which is what makes # this safe rather than the zeroing alone). if attention_mask is not None: attn_out = attn_out * attention_mask.view(batch_size, seq_len, 1).to(attn_out.dtype) # Output projection + residual out = self.out_proj(attn_out) out = self.dropout(out) return x + out class TransformerLayer(nn.Module): """Pre-RMSNorm transformer block: Attention + FFN.""" def __init__(self, config: DzairConfig) -> None: super().__init__() self.attention = Attention(config) self.ffn = FeedForward(config) def forward( self, x: Tensor, attention_mask: Tensor | None = None, is_causal: bool = False, ) -> Tensor: x = self.attention(x, attention_mask, is_causal) return self.ffn(x) class Embeddings(nn.Module): """Token embeddings with RMSNorm and dropout. No positional embeddings (RoPE handles position). Single-chunk inputs only: ``[CLS] chunk [SEP]``. """ def __init__(self, config: DzairConfig) -> None: super().__init__() self.word_embeddings = nn.Embedding( config.vocab_size, config.hidden_size, padding_idx=config.pad_token_id ) self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_eps) self.dropout = nn.Dropout(config.hidden_dropout_prob) def forward(self, input_ids: Tensor) -> Tensor: x = self.word_embeddings(input_ids) x = self.norm(x) return self.dropout(x) class Encoder(nn.Module): """Stack of transformer layers.""" def __init__(self, config: DzairConfig) -> None: super().__init__() self.config = config self.layers = nn.ModuleList( TransformerLayer(config) for _ in range(config.num_hidden_layers) ) self.gradient_checkpointing = False def forward( self, x: Tensor, attention_mask: Tensor | None = None, ) -> Tensor: for layer in self.layers: if self.gradient_checkpointing and self.training: x = checkpoint_utils.checkpoint( layer, x, attention_mask, False, use_reentrant=False ) else: x = layer(x, attention_mask, is_causal=False) # bidirectional return x def forward_with_states( self, x: Tensor, attention_mask: Tensor | None = None, ) -> tuple[Tensor, tuple[Tensor, ...]]: """Forward pass returning the last output plus each layer's output.""" states: list[Tensor] = [] for layer in self.layers: if self.gradient_checkpointing and self.training: x = checkpoint_utils.checkpoint( layer, x, attention_mask, False, use_reentrant=False ) else: x = layer(x, attention_mask, is_causal=False) # bidirectional states.append(x) return x, tuple(states) def set_gradient_checkpointing(model: nn.Module, value: bool) -> None: """Toggle activation checkpointing on every Encoder in a model. Plain attribute propagation, deliberately not via ``PreTrainedModel.gradient_checkpointing_enable`` whose signature drifted across transformers versions. The smoke test pins that outputs match and gradients flow with it on. """ for module in model.modules(): if isinstance(module, (Encoder, Generator)): module.gradient_checkpointing = value _ARCH_FIELDS: tuple[str, ...] = ( "vocab_size", "hidden_size", "intermediate_size", "num_attention_heads", "num_key_value_heads", "num_hidden_layers", "num_generator_layers", "generator_hidden_size", "generator_intermediate_size", "max_position_embeddings", "rope_theta", "hidden_dropout_prob", "attention_probs_dropout_prob", "layer_norm_eps", "pad_token_id", "cls_token_id", "sep_token_id", "mask_token_id", "tie_word_embeddings", "share_generator_embeddings", "qk_norm", ) _CONFIG_MISMATCH_MSG = ( "explicit config disagrees with the checkpoint's stored config on {field}: " "explicit={explicit!r} stored={stored!r} — pass config=None to trust the checkpoint" ) _NO_STORED_CONFIG_MSG = ( "checkpoint {path} carries no stored config and none was passed — " "pass config= explicitly" ) _FOLD_MISSING_MSG = ( "GDES checkpoint is missing {missing} — found prefixes: {prefixes}; " "cannot fold E_G + delta into the released embedding" ) _CHECKSUM_MISMATCH_MSG = "checkpoint checksum mismatch for {path}" def _normalize_stored_config(stored_dict: dict[str, Any]) -> dict[str, Any]: """Replace a pre-GQA null key-value count with the full-MHA default. Runs written before grouped-query attention store no (or null) key-value count; all of them trained full multi-head attention. """ normalized = dict(stored_dict) if normalized.get("num_key_value_heads") is None: normalized.pop("num_key_value_heads", None) return normalized def _check_config_match(config: DzairConfig, stored_dict: dict[str, Any]) -> None: """Raise on any recorded field the explicit config disagrees on. Fields the checkpoint predates (absent) or left null are not compared: the explicit config decides those, so era-appropriate explicit configs (global attention, full MHA, eval-only dropout) load instead of refusing on a formatting technicality. """ for field in _ARCH_FIELDS: if field not in stored_dict or stored_dict[field] is None: continue explicit_value = getattr(config, field, None) if explicit_value != stored_dict[field]: raise ValueError( _CONFIG_MISMATCH_MSG.format( field=field, explicit=explicit_value, stored=stored_dict[field] ) ) def _read_pretrain_checkpoint( checkpoint_path: str | Path, config: DzairConfig | None, ) -> tuple[DzairConfig, dict[str, Tensor]]: """Resolve (config, state) from a pretraining checkpoint. The checkpoint's stored config wins unless an explicit config is passed; an explicit config that disagrees with the stored one on a field the checkpoint actually records raises instead of silently misloading. Fields the checkpoint predates (absent) or left null are not compared: the explicit config decides those, so era-appropriate explicit configs (global attention, full MHA, eval-only dropout) load instead of refusing on a formatting technicality. Stored nulls/absences for the key-value count mean the run predates grouped-query attention and trained full multi-head attention, so they normalize to the default — never to a silent mismatch. """ path_obj = Path(checkpoint_path) sidecar = path_obj.parent / f"{path_obj.name}.sha256" if sidecar.is_file(): want = sidecar.read_text(encoding="utf-8").strip() digest = hashlib.sha256() with path_obj.open("rb") as f: for chunk in iter(lambda: f.read(1 << 20), b""): digest.update(chunk) if digest.hexdigest() != want: raise ValueError(_CHECKSUM_MISMATCH_MSG.format(path=checkpoint_path)) raw = torch.load(checkpoint_path, map_location="cpu", weights_only=True) if not isinstance(raw, dict): msg = f"checkpoint payload is not a mapping: {checkpoint_path}" raise TypeError(msg) inner = raw.get("model") state: dict[str, Tensor] = inner if isinstance(inner, dict) else raw stored = raw.get("config") stored_dict = stored if isinstance(stored, dict) else None if config is not None: if stored_dict is not None: _check_config_match(config, stored_dict) return copy.deepcopy(config), state if stored_dict is None: raise ValueError(_NO_STORED_CONFIG_MSG.format(path=checkpoint_path)) return DzairConfig(**_normalize_stored_config(stored_dict)), state def _fold_shared_backbone(state_dict: dict[str, Tensor]) -> dict[str, Tensor]: """Fold a GDES checkpoint's shared table into one released embedding. The released table is ``proj(E_G) + Δ`` — the generator's table through the width bridge plus the discriminator's delta — with the input norm taken from the discriminator's ``input_norm``. Same-width (or pre-bridge) checkpoints skip the projection, exactly like the forward does. """ out: dict[str, Tensor] = {} gen_key = "rtd_head.generator.embeddings.word_embeddings.weight" delta_key = "rtd_head.discriminator.delta_embeddings.weight" proj_key = "rtd_head.discriminator.gen_proj.weight" gen_table = state_dict.get(gen_key) delta = state_dict.get(delta_key) if gen_table is None or delta is None: missing = [k for k in (gen_key, delta_key) if k not in state_dict] prefixes = sorted({".".join(k.split(".")[:2]) if "." in k else k for k in state_dict}) raise KeyError(_FOLD_MISSING_MSG.format(missing=missing, prefixes=prefixes[:8])) proj = state_dict.get(proj_key) if proj is None or gen_table.size(-1) == delta.size(-1): folded = gen_table + delta.to(gen_table.dtype) else: folded = gen_table.to(proj.dtype) @ proj.T + delta.to(proj.dtype) out["embeddings.word_embeddings.weight"] = folded for key, value in state_dict.items(): if key.startswith("rtd_head.discriminator.input_norm."): out["embeddings.norm." + key[len("rtd_head.discriminator.input_norm.") :]] = value elif key.startswith("rtd_head.discriminator.encoder.") or key.startswith( "rtd_head.discriminator.norm." ): out[key[len("rtd_head.discriminator.") :]] = value return out _FUSED_SPLIT_MSG = ( "cannot map fused {key}: expected ({fused}, {hidden}), " "or the target is grouped-query ({kv} KV heads over {nq} query heads) " "which a fused full-MHA table cannot feed without lossy subsampling — " "retrain or load into a full-MHA config" ) _AMBIGUOUS_PROJ_MSG = ( "checkpoint mixes fused ({fused}) and split ({split}) attention projections — " "refusing instead of guessing which one owns the layer" ) _FUSED_BIAS_MSG = ( "cannot map fused {key}: biased projections have no split target — " "retrain or load into a matching config" ) def _unfuse_in_proj(state_dict: dict[str, Tensor], config: PretrainedConfig) -> dict[str, Tensor]: """Split fused full-MHA ``in_proj`` tables into ``q_proj`` + ``kv_proj``. Checkpoints written before grouped-query attention carry one fused QKV matrix per layer; current code keeps separate query and fused key/value projections. The split is exact only into full MHA (KV heads == query heads) with the canonical Q,K,V row order — anything else raises instead of silently remapping. Passes through states without fused tables untouched. """ fused_keys = [k for k in state_dict if k.endswith("attention.in_proj.weight")] if not fused_keys: return state_dict hidden = int(config.hidden_size) num_queries = int(config.num_attention_heads) num_kv = int(getattr(config, "num_key_value_heads", 0) or num_queries) split_keys = [ k for k in state_dict if k.endswith("attention.q_proj.weight") or k.endswith("attention.kv_proj.weight") ] if split_keys: msg = _AMBIGUOUS_PROJ_MSG.format(fused=fused_keys[0], split=split_keys[0]) raise RuntimeError(msg) biased = [k for k in state_dict if k.endswith("attention.in_proj.bias")] if biased: msg = _FUSED_BIAS_MSG.format(key=biased[0]) raise RuntimeError(msg) out = dict(state_dict) for key in fused_keys: weight = state_dict[key] if tuple(weight.shape) != (3 * hidden, hidden) or num_kv != num_queries: msg = _FUSED_SPLIT_MSG.format( key=key, fused=tuple(weight.shape), hidden=hidden, kv=num_kv, nq=num_queries, ) raise RuntimeError(msg) prefix = key[: -len("in_proj.weight")] query, key_p, value = weight.split([hidden, hidden, hidden], dim=0) del out[key] out[prefix + "q_proj.weight"] = query out[prefix + "kv_proj.weight"] = torch.cat([key_p, value], dim=0) return out _GEN_TABLE_KEY = "rtd_head.generator.embeddings.word_embeddings.weight" def discriminator_backbone_state( state_dict: dict[str, Tensor], config: PretrainedConfig ) -> dict[str, Tensor]: """Map a pretraining checkpoint's discriminator weights onto ``DzairModel``. GDES checkpoints (shared table): the released embedding is the fold ``proj(E_G) + Δ`` — the generator's table through the width bridge plus the discriminator's delta — with the input norm taken from the discriminator's ``input_norm``. Independent checkpoints: ``rtd_head.discriminator.*`` maps verbatim minus the RTD classifier. A payload that is already a ``DzairModel`` state dict (no ``rtd_head`` prefix) passes through; ``strict=True`` on the caller's ``load_state_dict`` catches anything malformed. Fused full-MHA ``in_proj`` tables are split exactly (see ``_unfuse_in_proj``); generator-trunk keys never enter the mapping, so a fused generator neither helps nor breaks the fold. """ shared = bool(getattr(config, "share_generator_embeddings", False)) if not any(k.startswith("rtd_head.") for k in state_dict): return { (key[len("dzair.") :] if key.startswith("dzair.") else key): value for key, value in state_dict.items() } relevant = { key: value for key, value in state_dict.items() if key.startswith("rtd_head.discriminator.") or key == _GEN_TABLE_KEY or key.startswith("dzair.") } state_dict = _unfuse_in_proj(relevant, config) out: dict[str, Tensor] = {} if shared: return _fold_shared_backbone(state_dict) for key, value in state_dict.items(): if key.startswith("rtd_head.discriminator.") and not key.startswith( "rtd_head.discriminator.classifier" ): out[key[len("rtd_head.discriminator.") :]] = value elif key.startswith("dzair."): out[key[len("dzair.") :]] = value return out # Pretraining-only modules absent from older checkpoints: a checkpoint missing # exactly these still loads, everything else missing or unexpected still raises. _COMPAT_MISSING_SUBSTRINGS: tuple[str, ...] = ("gen_proj.",) _GENERATION_GAP_MSG = ( "checkpoint uses independent discriminator embeddings " "('rtd_head.discriminator.embeddings.') but the model expects GDES " "('rtd_head.discriminator.delta_embeddings.'): no automatic migration — " "the v1 identity (E_D independent) cannot fold into E_G + delta without " "changing numerics; retrain or load into a share_generator_embeddings=False " "config" ) def load_pretrain_state(model: nn.Module, state: dict[str, Tensor]) -> None: """Load a pretraining state dict across the width-bridge generation gap. Checkpoints written before the generator width bridge lack ``gen_proj``; anything else missing, misshapen, or unexpected still raises. Fused full-MHA ``in_proj`` tables are split exactly (see ``_unfuse_in_proj``). The independent-embeddings (v1) to GDES generation gap is refused loudly: silently mapping E_D onto delta would change numerics. """ model_config = getattr(model, "config", None) if model_config is not None: gen_keys = {k: v for k, v in state.items() if k.startswith("rtd_head.generator.encoder.")} trunk_keys = {k: v for k, v in state.items() if k not in gen_keys} merged_state = _unfuse_in_proj(trunk_keys, model_config) if gen_keys: merged_state.update(_unfuse_in_proj(gen_keys, _generator_view(model_config))) state = merged_state own = model.state_dict() if any("rtd_head.discriminator.embeddings." in k for k in state) and any( "delta_embeddings" in k for k in own ): raise RuntimeError(_GENERATION_GAP_MSG) if any("delta_embeddings" in k for k in state) and any( "rtd_head.discriminator.embeddings." in k for k in own ): raise RuntimeError(_GENERATION_GAP_MSG) unexpected = [k for k in state if k not in own] if unexpected: msg = f"checkpoint holds unexpected keys: {unexpected[:8]}" raise RuntimeError(msg) merged: dict[str, Tensor] = {} absent: list[str] = [] for key, value in own.items(): if key not in state: absent.append(key) continue if value.shape != state[key].shape: msg = f"checkpoint shape mismatch for {key}: ckpt {tuple(state[key].shape)}" raise RuntimeError(msg) merged[key] = state[key] unaccounted = [k for k in absent if not any(s in k for s in _COMPAT_MISSING_SUBSTRINGS)] if unaccounted: msg = f"checkpoint lacks load-bearing keys: {unaccounted[:8]}" raise RuntimeError(msg) for key in absent: merged[key] = own[key] model.load_state_dict(merged, strict=True) # State keys from retired training objectives. A checkpoint carrying them # predates the current code: the trunk weights still load, the retired # heads do not come back. Centralized here so every loader agrees on # what "obsolete" means; anything else unexpected still raises. OBSOLETE_STATE_SUBSTRINGS: tuple[str, ...] = ( "order_head.", "order_loss_ema", "token_loss_ema", "token_type_embeddings.", ) @dataclass(frozen=True) class ResumeCompat: """How a checkpoint's weights mapped onto the current model.""" generation: str # "same" (exact) or "legacy" (obsolete keys dropped) dropped: tuple[str, ...] def load_resume_weights(model: nn.Module, ckpt_model_state: dict[str, Tensor]) -> ResumeCompat: """Load training weights for an exact resume across code generations. Fused full-MHA tables split exactly (trunk and generator widths handled separately); retired keys drop loudly in the report. Any other missing, misshapen, or unexpected key raises — a half-mapped model never trains. The caller decides from ``generation`` whether the optimizer may be restored (``same``) or must restart fresh (``legacy``): stale momentum on a reshaped model is silent corruption. """ raw = {k.removeprefix("_orig_mod."): v for k, v in ckpt_model_state.items()} model_config = getattr(model, "config", None) if model_config is not None: gen_keys = {k: v for k, v in raw.items() if k.startswith("rtd_head.generator.encoder.")} trunk_keys = {k: v for k, v in raw.items() if k not in gen_keys} raw = _unfuse_in_proj(trunk_keys, model_config) if gen_keys: raw.update(_unfuse_in_proj(gen_keys, _generator_view(model_config))) dropped = tuple(sorted({k for k in raw if any(s in k for s in OBSOLETE_STATE_SUBSTRINGS)})) kept = {k: v for k, v in raw.items() if k not in dropped} raw_model = getattr(model, "_orig_mod", model) own = raw_model.state_dict() unexpected = [k for k in kept if k not in own] if unexpected: msg = f"checkpoint holds unexpected keys: {unexpected[:8]}" raise RuntimeError(msg) missing = [k for k in own if k not in kept] if missing: msg = f"checkpoint lacks load-bearing keys: {missing[:8]}" raise RuntimeError(msg) for key, value in own.items(): if value.shape != kept[key].shape: msg = f"checkpoint shape mismatch for {key}: ckpt {tuple(kept[key].shape)}" raise RuntimeError(msg) raw_model.load_state_dict(kept, strict=True) return ResumeCompat(generation="legacy" if dropped else "same", dropped=dropped) def _generator_view(config: DzairConfig) -> DzairConfig: """A config view sizing the generator trunk: narrow width, global attention. The generator keeps head_dim 64 and key-value groups proportional to the trunk; it always attends globally so corruption quality never depends on the discriminator's local window. Copies (never mutates) the trunk config. """ view = copy.copy(config) view.hidden_size = config.generator_hidden_size view.intermediate_size = config.generator_intermediate_size view.num_attention_heads = config.generator_num_heads view.num_key_value_heads = max( 1, config.generator_num_heads * config.num_key_value_heads // config.num_attention_heads ) if view.num_attention_heads % view.num_key_value_heads != 0: msg = ( f"generator {view.num_attention_heads} query heads must split over " f"{view.num_key_value_heads} key-value heads" ) raise ValueError(msg) view.num_hidden_layers = config.num_generator_layers return view class Generator(nn.Module): """Lightweight MLM generator for RTD corruption. GDES: embeddings shared, detached for discriminator. The generator trunk runs at ``generator_hidden_size`` behind a width projection only where it meets the discriminator (see ``Discriminator.gen_proj``); its own input and LM head stay in the narrow width with tied tables. """ def __init__(self, config: DzairConfig) -> None: super().__init__() self.config = config view = _generator_view(config) self.embeddings = Embeddings(view) self.encoder = nn.ModuleList( TransformerLayer(view) for _ in range(config.num_generator_layers) ) self.norm = RMSNorm(view.hidden_size, eps=config.layer_norm_eps) self.lm_head = nn.Linear(view.hidden_size, config.vocab_size, bias=False) # Tie output embeddings to input embeddings self.lm_head.weight = self.embeddings.word_embeddings.weight self.gradient_checkpointing = False def forward( self, input_ids: Tensor, attention_mask: Tensor | None = None, ) -> Tensor: x = self.embeddings(input_ids) for layer in self.encoder: if self.gradient_checkpointing and self.training: x = checkpoint_utils.checkpoint( layer, x, attention_mask, False, use_reentrant=False ) else: x = layer(x, attention_mask, is_causal=False) x = self.norm(x) return self.lm_head(x) class Discriminator(nn.Module): """RTD discriminator: detects replaced tokens. Two embedding policies, selected by ``config.share_generator_embeddings``: - **GDES** (True): the discriminator reads ``proj(stop_grad(E_G)) + Δ`` where ``E_G`` is the generator's own (narrow) table — generator MLM training shapes the table the discriminator reads — ``proj`` bridges the generator width to the trunk width, and ``Δ`` is this module's own table. Discriminator gradients flow to ``Δ`` and ``proj`` only, by construction. The released backbone folds ``proj(E_G) + Δ`` into one table at load time. - **Independent** (False, default): a private ``Embeddings`` table, as in classic ELECTRA. The pretraining head still passes the generator's table; in this mode it is unused, and the forward is a plain lookup. """ def __init__(self, config: DzairConfig) -> None: super().__init__() self.config = config self.share_generator_embeddings = bool(config.share_generator_embeddings) if self.share_generator_embeddings: self.delta_embeddings = nn.Embedding( config.vocab_size, config.hidden_size, padding_idx=config.pad_token_id ) self.gen_proj = nn.Linear(config.generator_hidden_size, config.hidden_size, bias=False) self.input_norm = RMSNorm(config.hidden_size, eps=config.layer_norm_eps) self.input_dropout = nn.Dropout(config.hidden_dropout_prob) else: self.embeddings = Embeddings(config) self.encoder = Encoder(config) self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_eps) self.classifier = nn.Linear(config.hidden_size, 2, bias=False) # binary: original/replaced def forward( self, input_ids: Tensor, attention_mask: Tensor | None = None, generator_embeddings: Tensor | None = None, ) -> Tensor: """``generator_embeddings`` is the generator's table (required under GDES).""" if self.share_generator_embeddings: if generator_embeddings is None: msg = "share_generator_embeddings=True requires the generator table" raise ValueError(msg) base = functional.embedding( input_ids, generator_embeddings.detach(), padding_idx=self.config.pad_token_id ) if base.size(-1) != self.config.hidden_size: base = self.gen_proj(base) x = base + self.delta_embeddings(input_ids) x = self.input_dropout(self.input_norm(x)) else: x = self.embeddings(input_ids) x = self.encoder(x, attention_mask) x = self.norm(x) return self.classifier(x) def _dynamo_disabled[FnT](function: FnT) -> FnT: """Exclude CPU-scalar bookkeeping from the compiled graph. ``float()`` syncs inside the forward break Dynamo (measured: a ``Tensor.item()`` graph break every step) and stall the GPU for a value only needed as an eager loss scale. Falls back to a no-op when ``disable`` is unavailable; the pinned images always carry it. """ disable = getattr(_dynamo, "disable", None) if disable is None: return function return disable(function) @_dynamo_disabled def draw_token_mask(candidate: Tensor, mask_prob: float | Tensor) -> Tensor: """Per-token uniform draw over candidate positions.""" return candidate & (torch.rand(candidate.shape, device=candidate.device) < mask_prob) @_dynamo_disabled def draw_word_mask(candidate: Tensor, word_starts: Tensor, mask_prob: float | Tensor) -> Tensor: """One uniform draw per word; every candidate position in a chosen word masked. Positions before the first word start are never masked. """ device = candidate.device length = candidate.size(-1) arange = torch.arange(length, device=device).expand_as(candidate) cur_start = torch.where(word_starts, arange, -1).cummax(dim=-1).values chosen = word_starts & candidate & (torch.rand(candidate.shape, device=device) < mask_prob) last_chosen = torch.where(chosen, arange, -1).cummax(dim=-1).values return candidate & (cur_start >= 0) & (cur_start == last_chosen) @dataclass(frozen=True) class MaskSpec: """What may be masked and how often (built per step from the schedule).""" special_ids: frozenset[int] vocab_size: int mask_token_id: int mask_prob: float | Tensor @_dynamo_disabled def _mask_inputs( input_ids: Tensor, eligible: Tensor, spec: MaskSpec, word_starts: Tensor | None = None, ) -> tuple[Tensor, Tensor]: """BERT 80/10/10 corruption. Returns (masked_input_ids, mlm_labels). Masking is whole-word when ``word_starts`` ([B, T] bool, True at word-initial pieces) is given, else per-token. Special ids (pad/cls/sep) and ineligible positions are never masked. Dynamic every step. """ device = input_ids.device is_special = torch.zeros_like(input_ids, dtype=torch.bool) for sid in spec.special_ids: is_special |= input_ids == sid candidate = eligible & ~is_special if word_starts is not None: masked = draw_word_mask(candidate, word_starts, spec.mask_prob) else: masked = draw_token_mask(candidate, spec.mask_prob) rand = torch.rand(input_ids.shape, device=device) replace_mask = masked & (rand < MASK_REPLACE_CUTOFF) random_mask = masked & (rand >= MASK_REPLACE_CUTOFF) & (rand < MASK_RANDOM_CUTOFF) # keep_mask (last 10%): input unchanged, still predicted. masked_input = input_ids.clone() masked_input[replace_mask] = spec.mask_token_id rand_tokens = torch.randint_like(input_ids, 0, spec.vocab_size) masked_input = torch.where(random_mask, rand_tokens, masked_input) mlm_labels = torch.full_like(input_ids, IGNORE_INDEX) mlm_labels[masked] = input_ids[masked] return masked_input, mlm_labels @dataclass class DzairRTDOutput: """Pretraining output: ELECTRA-style joint loss.""" loss: Tensor | None rtd_logits: Tensor gen_logits: Tensor generator_loss: Tensor | None discriminator_loss: Tensor | None replacement_rate: Tensor @_dynamo_disabled def _sample_generator_corruptions( gen_logits: Tensor, masked_input: Tensor, predict: Tensor, input_ids: Tensor, softmax_chunk: int = 2048, ) -> Tensor: """Sample generator tokens on masked positions in chunks outside Dynamo.""" with torch.no_grad(): flat_mask = predict.reshape(-1) idx = torch.where(flat_mask)[0] sampled = torch.empty_like(idx) flat_logits = gen_logits.reshape(-1, gen_logits.size(-1)) for start in range(0, idx.numel(), softmax_chunk): group = idx[start : start + softmax_chunk] probs = flat_logits[group].float().softmax(dim=-1) sampled[start : start + softmax_chunk] = torch.multinomial(probs, 1).squeeze(-1) corrupted = masked_input.reshape(-1).clone() corrupted[idx] = sampled return corrupted.view_as(input_ids) class RTDHead(nn.Module): """Generator (MLM) corrupts, discriminator (RTD) detects, GDES detaches.""" def __init__(self, config: DzairConfig, rtd_loss_weight: float = RTD_LOSS_WEIGHT) -> None: super().__init__() self.config = config self.rtd_loss_weight = rtd_loss_weight self.generator = Generator(config) self.discriminator = Discriminator(config) def forward( self, input_ids: Tensor, attention_mask: Tensor | None = None, mask_prob: float | Tensor = 0.15, word_starts: Tensor | None = None, ) -> DzairRTDOutput: """Returns the joint output. ``mask_prob`` follows the 30→15% schedule. Accepts a 0-dim tensor as well as a float: pass a tensor from any compiled caller — Dynamo specializes on float argument *values*, so a per-step float schedule would recompile every step until the cache limit forces the whole model back to eager. """ eligible = ( attention_mask.to(torch.bool) if attention_mask is not None else torch.ones_like(input_ids, dtype=torch.bool) ) spec = MaskSpec( special_ids=frozenset( sid for sid in ( self.config.pad_token_id, self.config.cls_token_id, self.config.sep_token_id, ) if sid is not None ), vocab_size=self.config.vocab_size, mask_token_id=self.config.mask_token_id, mask_prob=mask_prob, ) masked_input, mlm_labels = _mask_inputs(input_ids, eligible, spec, word_starts) gen_logits = self.generator(masked_input, attention_mask) gen_loss: Tensor | None = None predict = mlm_labels != IGNORE_INDEX if predict.any(): gen_loss = functional.cross_entropy(gen_logits[predict], input_ids[predict].detach()) # Corrupt only the masked positions by sampling the generator. # Softmax runs over the masked subset in bounded chunks to cap the # peak transient allocation. Arithmetic (measured 2026-09-10): # 8192 * 48000 * 4 bytes = 1.57 GB -- OOMs at 20.94 GB in use # 2048 * 48000 * 4 bytes = 0.39 GB -- 4.7 GB headroom at 18.9 GB peak # Chunked multinomial is mathematically identical to sampling all at once. _softmax_chunk = 2048 corrupted = _sample_generator_corruptions( gen_logits, masked_input, predict, input_ids, _softmax_chunk ) disc_logits = self.discriminator( corrupted, attention_mask, generator_embeddings=self.generator.embeddings.word_embeddings.weight, ) rtd_labels = torch.where(corrupted == input_ids, 1, 0) rtd_labels = torch.where(eligible, rtd_labels, IGNORE_INDEX) disc_loss: Tensor | None = None if (rtd_labels != IGNORE_INDEX).any(): disc_loss = functional.cross_entropy( disc_logits.reshape(-1, 2), rtd_labels.reshape(-1), ignore_index=IGNORE_INDEX ) loss: Tensor | None = None if gen_loss is not None and disc_loss is not None: loss = gen_loss + self.rtd_loss_weight * disc_loss with torch.no_grad(): replacement_rate = ( (corrupted[predict] != input_ids[predict]).float().mean() if predict.any() else torch.zeros((), device=input_ids.device) ) return DzairRTDOutput( loss=loss, rtd_logits=disc_logits, gen_logits=gen_logits, generator_loss=gen_loss, discriminator_loss=disc_loss, replacement_rate=replacement_rate, ) class DzairPreTrainedModel(PreTrainedModel): config_class = DzairConfig base_model_prefix = "dzair" supports_gradient_checkpointing = True _no_split_modules: ClassVar[list[str]] = ["TransformerLayer"] def _init_weights(self, module: nn.Module) -> None: # Masinissa scaled init, depth-scaled trunc normal; deliberately not # config-driven (a field that silently does nothing is worse than none). std = math.sqrt(2.0 / (5.0 * self.config.hidden_size)) if isinstance(module, nn.Linear): nn.init.trunc_normal_(module.weight, mean=0.0, std=std, a=-2 * std, b=2 * std) if module.bias is not None: nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): nn.init.trunc_normal_(module.weight, mean=0.0, std=std, a=-2 * std, b=2 * std) elif isinstance(module, RMSNorm): nn.init.ones_(module.weight) @dataclass class DzairEncoderOutput(ModelOutput): """Encoder output: last state plus optional per-layer states. A dedicated type because the framework's BaseModelOutput pins its state fields to FloatTensor, which the checker treats as distinct from Tensor. hidden_states[0] is the embedding output (HF convention). Per-layer entries are pre-norm layer outputs; last_hidden_state is post-norm, so hidden_states[-1] != last_hidden_state by design. """ last_hidden_state: Tensor hidden_states: tuple[Tensor, ...] | None = None attentions: tuple[Tensor, ...] | None = None class DzairModel(DzairPreTrainedModel): """The encoder alone (discriminator backbone). Returns contextualised token representations. """ def __init__(self, config: DzairConfig) -> None: super().__init__(config) self.embeddings = Embeddings(config) self.encoder = Encoder(config) self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_eps) self.post_init() def get_input_embeddings(self) -> nn.Embedding: return self.embeddings.word_embeddings def set_input_embeddings(self, value: nn.Embedding) -> None: self.embeddings.word_embeddings = value def forward( self, input_ids: Tensor, attention_mask: Tensor | None = None, output_hidden_states: bool = False, ) -> DzairEncoderOutput: """Encode tokens. With output_hidden_states, hidden_states[0] is the embedding output and [k] the k-th layer output (HF convention). Layer states are pre-norm; last_hidden_state is post-norm. """ embedded = self.embeddings(input_ids) if output_hidden_states: last, states = self.encoder.forward_with_states(embedded, attention_mask) return DzairEncoderOutput( last_hidden_state=self.norm(last), hidden_states=(embedded, *states), ) x = self.encoder(embedded, attention_mask) return DzairEncoderOutput(last_hidden_state=self.norm(x)) class DzairForMaskedLM(DzairPreTrainedModel): """Pretraining model: Generator (MLM) + Discriminator (RTD) with GDES.""" _tied_weights_keys: ClassVar[dict[str, str]] = { "rtd_head.generator.lm_head.weight": "rtd_head.generator.embeddings.word_embeddings.weight", } def __init__(self, config: DzairConfig) -> None: super().__init__(config) self.rtd_head = RTDHead(config) self.post_init() def forward( self, input_ids: Tensor, attention_mask: Tensor | None = None, mask_prob: float | Tensor = 0.15, word_starts: Tensor | None = None, ) -> DzairRTDOutput: return self.rtd_head(input_ids, attention_mask, mask_prob, word_starts) @dataclass class DzairSequenceClassifierOutput(ModelOutput): """Output type of DzairForSequenceClassification.""" loss: Tensor | None = None logits: Tensor | None = None hidden_states: tuple[Tensor, ...] | None = None attentions: tuple[Tensor, ...] | None = None class DzairForSequenceClassification(DzairPreTrainedModel): """Sequence classification head on top of the DZAIR encoder backbone. One method, the measured one: [CLS] pooling through an MLP projection head (Dropout -> Dense -> GELU -> Dropout) into the classification layer. The DZNLI head ablation picked cls+mlp over mean+linear; the landmark and attention experiments never measured a win, so they do not ship. """ def __init__(self, config: DzairConfig) -> None: super().__init__(config) self.num_labels = getattr(config, "num_labels", 2) self.dzair = DzairModel(config) self.dropout = nn.Dropout(config.hidden_dropout_prob) self.dense = nn.Linear(config.hidden_size, config.hidden_size) self.classifier = nn.Linear(config.hidden_size, self.num_labels) self.post_init() def get_input_embeddings(self) -> nn.Embedding: return self.dzair.get_input_embeddings() def set_input_embeddings(self, value: nn.Embedding) -> None: self.dzair.set_input_embeddings(value) def forward( self, input_ids: Tensor, attention_mask: Tensor | None = None, labels: Tensor | None = None, output_hidden_states: bool = False, ) -> DzairSequenceClassifierOutput: outputs = self.dzair( input_ids, attention_mask=attention_mask, output_hidden_states=output_hidden_states ) pooled_output = self.dropout(outputs.last_hidden_state[:, 0]) pooled_output = self.dense(pooled_output) pooled_output = functional.gelu(pooled_output) pooled_output = self.dropout(pooled_output) logits = self.classifier(pooled_output) loss: Tensor | None = None if labels is not None: if self.num_labels == 1: loss = functional.mse_loss(logits.view(-1), labels.view(-1).float()) else: loss = functional.cross_entropy(logits.view(-1, self.num_labels), labels.view(-1)) return DzairSequenceClassifierOutput( loss=loss, logits=logits, hidden_states=outputs.hidden_states, attentions=outputs.attentions, ) def load_backbone_weights(self, state_dict: dict[str, Tensor]) -> None: """Load pretrained discriminator backbone weights into self.dzair.""" self.dzair.load_state_dict( discriminator_backbone_state(state_dict, self.config), strict=True ) @classmethod def from_pretrained_checkpoint( cls, checkpoint_path: str | Path, config: DzairConfig | None = None, num_labels: int = 2, ) -> DzairForSequenceClassification: """Instantiate classification model and load backbone from pretrain checkpoint. The checkpoint's stored config is preferred; an explicit config that disagrees with it raises (see ``_read_pretrain_checkpoint``). """ model_config, state = _read_pretrain_checkpoint(checkpoint_path, config) model_config.num_labels = num_labels model = cls(model_config) model.load_backbone_weights(state) return model @dataclass class DzairTokenClassifierOutput(ModelOutput): """Output type of DzairForTokenClassification.""" loss: Tensor | None = None logits: Tensor | None = None hidden_states: tuple[Tensor, ...] | None = None attentions: tuple[Tensor, ...] | None = None class DzairForTokenClassification(DzairPreTrainedModel): """Token classification head on top of the DZAIR encoder backbone (e.g. for NER/POS).""" def __init__(self, config: DzairConfig) -> None: super().__init__(config) self.num_labels = getattr(config, "num_labels", 2) self.dzair = DzairModel(config) self.dropout = nn.Dropout(config.hidden_dropout_prob) self.classifier = nn.Linear(config.hidden_size, self.num_labels) self.post_init() def get_input_embeddings(self) -> nn.Embedding: return self.dzair.get_input_embeddings() def set_input_embeddings(self, value: nn.Embedding) -> None: self.dzair.set_input_embeddings(value) def forward( self, input_ids: Tensor, attention_mask: Tensor | None = None, labels: Tensor | None = None, ) -> DzairTokenClassifierOutput: outputs = self.dzair(input_ids, attention_mask=attention_mask) sequence_output = outputs.last_hidden_state sequence_output = self.dropout(sequence_output) logits = self.classifier(sequence_output) loss: Tensor | None = None if labels is not None: loss = functional.cross_entropy( logits.view(-1, self.num_labels), labels.view(-1), ignore_index=-100, ) return DzairTokenClassifierOutput( loss=loss, logits=logits, hidden_states=outputs.hidden_states, attentions=outputs.attentions, ) def load_backbone_weights(self, state_dict: dict[str, Tensor]) -> None: """Load pretrained discriminator backbone weights into self.dzair.""" self.dzair.load_state_dict( discriminator_backbone_state(state_dict, self.config), strict=True ) @classmethod def from_pretrained_checkpoint( cls, checkpoint_path: str | Path, config: DzairConfig | None = None, num_labels: int = 2, ) -> DzairForTokenClassification: """Instantiate token classification model and load backbone from pretrain checkpoint.""" model_config, state = _read_pretrain_checkpoint(checkpoint_path, config) model_config.num_labels = num_labels model = cls(model_config) model.load_backbone_weights(state) return model __all__ = [ "OBSOLETE_STATE_SUBSTRINGS", "RTD_LOSS_WEIGHT", "DzairConfig", "DzairEncoderOutput", "DzairForMaskedLM", "DzairForSequenceClassification", "DzairForTokenClassification", "DzairModel", "DzairPreTrainedModel", "DzairRTDOutput", "DzairSequenceClassifierOutput", "DzairTokenClassifierOutput", "MaskSpec", "ResumeCompat", "discriminator_backbone_state", "draw_token_mask", "draw_word_mask", "load_pretrain_state", "load_resume_weights", "set_gradient_checkpointing", ]