File size: 2,147 Bytes
bea429c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
"""Load the exact trained Gemma scorer and expose its scalar logit path."""

from __future__ import annotations

from assets import LOCK, fetch_base, fetch_source


def load_trained_scorer():
    """Return tokenizer and merged scorer, refusing an absent or altered trained head."""
    # Access check happens before any model allocation or download of large base weights.
    base = fetch_base()
    source = fetch_source()

    import torch
    from peft import PeftModel
    from safetensors.torch import load_file
    from transformers import AutoTokenizer, Gemma3TextForSequenceClassification

    tokenizer = AutoTokenizer.from_pretrained(base, local_files_only=True)
    if tokenizer.pad_token is None:
        tokenizer.pad_token = tokenizer.eos_token
    model = Gemma3TextForSequenceClassification.from_pretrained(
        base, num_labels=1, dtype=torch.float32, local_files_only=True
    )
    model.config.pad_token_id = tokenizer.pad_token_id
    model.config.eos_token_id = tokenizer.eos_token_id
    peft = PeftModel.from_pretrained(model, source / "pretrained-scorer", local_files_only=True)
    merged = peft.merge_and_unload().eval()
    trained = load_file(source / "pretrained-scorer" / "adapter_model.safetensors")
    trained_head = trained[LOCK["native_serving"]["trained_score_tensor"]].float()
    actual_head = merged.score.weight.detach().float()
    if tuple(trained_head.shape) != tuple(LOCK["native_serving"]["score_shape"]):
        raise ValueError("locked trained head shape changed")
    if not torch.equal(trained_head, actual_head):
        raise ValueError("merged model lost the trained scalar score head")
    return tokenizer, merged


def export_wrapper(model):
    """Make a traceable wrapper that returns only the trained choice logit."""
    import torch

    class ScalarScorer(torch.nn.Module):
        def __init__(self, native):
            super().__init__()
            self.native = native

        def forward(self, input_ids, attention_mask):
            return self.native(input_ids=input_ids.long(), attention_mask=attention_mask.long()).logits.float()

    return ScalarScorer(model).eval()