File size: 35,976 Bytes
fa5ff8a
 
2a268ef
 
 
 
 
 
 
 
 
 
 
 
fa5ff8a
 
2a268ef
fa5ff8a
 
 
2a268ef
 
 
 
 
 
741fa39
 
 
 
2a268ef
a71717d
fa5ff8a
 
 
 
 
 
 
a71717d
 
 
 
 
 
 
 
 
 
 
 
fa5ff8a
2a268ef
 
 
fa5ff8a
2a268ef
 
 
 
fa5ff8a
 
 
2a268ef
 
 
 
 
 
 
 
fa5ff8a
 
 
2a268ef
fa5ff8a
 
 
 
 
 
 
 
 
 
 
 
 
 
2a268ef
fa5ff8a
 
 
 
 
 
2a268ef
 
 
 
 
 
 
fa5ff8a
 
 
2a268ef
fa5ff8a
 
 
 
2a268ef
fa5ff8a
 
2a268ef
741fa39
2a268ef
fa5ff8a
 
741fa39
 
2a268ef
 
fa5ff8a
 
 
 
 
2a268ef
 
 
 
 
 
 
fa5ff8a
 
a71717d
741fa39
fa5ff8a
 
 
 
2a268ef
 
 
741fa39
2a268ef
 
741fa39
2a268ef
fa5ff8a
2a268ef
 
 
 
 
 
 
 
 
fa5ff8a
2a268ef
fa5ff8a
741fa39
 
fa5ff8a
 
 
 
 
 
 
 
2a268ef
 
 
 
 
 
 
fa5ff8a
 
 
 
 
 
2a268ef
fa5ff8a
 
 
 
 
 
 
 
 
 
 
 
 
b94233a
fa5ff8a
 
 
 
 
2a268ef
b94233a
 
 
fa5ff8a
2a268ef
fa5ff8a
 
 
 
 
2a268ef
fa5ff8a
 
 
2a268ef
 
fa5ff8a
 
 
 
 
 
 
2a268ef
 
 
 
fa5ff8a
 
 
 
 
2a268ef
 
fa5ff8a
2a268ef
fa5ff8a
 
2a268ef
 
 
 
 
 
 
fa5ff8a
 
a71717d
 
 
 
2a268ef
a71717d
fa5ff8a
 
2a268ef
fa5ff8a
 
 
 
 
 
2a268ef
 
fa5ff8a
 
 
 
 
 
 
 
2a268ef
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fa5ff8a
 
2a268ef
741fa39
fa5ff8a
 
741fa39
fa5ff8a
 
 
 
741fa39
fa5ff8a
 
 
 
 
 
2a268ef
fa5ff8a
2a268ef
741fa39
2a268ef
fa5ff8a
 
741fa39
2a268ef
741fa39
 
fa5ff8a
 
 
 
 
 
 
 
 
 
 
 
2a268ef
fa5ff8a
 
2a268ef
fa5ff8a
741fa39
fa5ff8a
2a268ef
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fa5ff8a
2a268ef
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fa5ff8a
2a268ef
 
 
fa5ff8a
 
2a268ef
 
741fa39
 
2a268ef
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fa5ff8a
 
 
 
 
 
 
741fa39
 
 
 
2a268ef
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
741fa39
fa5ff8a
 
741fa39
 
 
fa5ff8a
2a268ef
 
 
 
 
 
 
 
 
 
 
fa5ff8a
 
2a268ef
 
 
fa5ff8a
2a268ef
 
 
 
fa5ff8a
2a268ef
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
import torch
from torch import nn
from peft import (
    get_peft_model,
    LoraConfig,
    TaskType,
    AutoPeftModelForCausalLM,
    AutoPeftModelForSequenceClassification,
)
from transformers import (
    AutoModelForCausalLM,
    AutoModelForSequenceClassification,
    AutoTokenizer,
)
import time
import json
import random

import os

try:
    from transformers import AdamW
except ImportError:
    from torch.optim import AdamW


def calculate_MMD_loss(human_crit, sample_crit):
    mmd_loss = human_crit.mean() - sample_crit.mean()
    return mmd_loss


def from_pretrained(cls, model_name, kwargs, cache_dir, device=None):
    # use local model if it exists
    if "/" in model_name:
        local_path = os.path.join(cache_dir, model_name.split("/")[1])
    else:
        local_path = os.path.join(cache_dir, model_name)

    if os.path.exists(local_path):
        return cls.from_pretrained(local_path, **kwargs, trust_remote_code=True)

    remote_kwargs = dict(kwargs, cache_dir=cache_dir, trust_remote_code=True)
    if device is not None:
        # Pin the whole model to a single device instead of device_map='auto'.
        # 'auto' lets accelerate split the model across GPU/CPU/disk when
        # memory is tight at load time, which then makes a later `.to(device)`
        # raise "You can't move a model that has some modules offloaded to
        # cpu or disk." Forcing everything onto one device up front avoids
        # that split entirely.
        remote_kwargs["device_map"] = {"": device}
    return cls.from_pretrained(model_name, **remote_kwargs)


model_fullnames = {
    # default base for the witness functions
    'gemma-1b': 'google/gemma-3-1b-pt',
    'gemma-4b': 'google/gemma-3-4b-pt',
    'qwen-1.5b': 'Qwen/Qwen2.5-1.5B',
    'falcon-1b': 'tiiuae/Falcon3-1B-Base',
    'phi-1b': 'microsoft/phi-1',
}
float16_models = []

# Dtype for the *trainable* scoring model under gemma-1b. fp32 keeps the LoRA
# weight updates exact (bf16 would round away small lr*grad steps); the frozen
# reference model is always bf16. Switch this to torch.bfloat16 to also halve
# the scoring model's weight memory if you still hit OOM, at a small risk to
# AdaJASA training precision.
GEMMA1B_SCORING_DTYPE = torch.float32


def get_model_fullname(model_name):
    return model_fullnames[model_name] if model_name in model_fullnames else model_name


def load_tokenizer(model_name, for_dataset, cache_dir):
    model_fullname = get_model_fullname(model_name)
    optional_tok_kwargs = {}
    if for_dataset in ['pubmed']:
        optional_tok_kwargs['padding_side'] = 'left'
    else:
        optional_tok_kwargs['padding_side'] = 'right'
    base_tokenizer = from_pretrained(AutoTokenizer, model_fullname, optional_tok_kwargs, cache_dir=cache_dir)
    if base_tokenizer.pad_token_id is None:
        base_tokenizer.pad_token_id = base_tokenizer.eos_token_id
        if '13b' in model_fullname:
            base_tokenizer.pad_token_id = 0
    return base_tokenizer


def get_sampling_discrepancy_analytic(logits_ref, logits_score, labels):
    if logits_ref.size(-1) != logits_score.size(-1):
        vocab_size = min(logits_ref.size(-1), logits_score.size(-1))
        logits_ref = logits_ref[:, :, :vocab_size]
        logits_score = logits_score[:, :, :vocab_size]

    # Evaluate the witness statistic in fp32 even when the models run in bf16:
    # the full-vocab softmax, the probability-weighted variance and the
    # var-normalised division are precision-sensitive. The upcast buffer is
    # transient (freed each step), so memory impact is small.
    logits_ref = logits_ref.float()
    logits_score = logits_score.float()

    labels = labels.unsqueeze(-1) if labels.ndim == logits_score.ndim - 1 else labels
    lprobs_score = torch.log_softmax(logits_score, dim=-1)
    probs_ref = torch.softmax(logits_ref, dim=-1)

    log_likelihood = lprobs_score.gather(dim=-1, index=labels).squeeze(-1)
    mean_ref = (probs_ref * lprobs_score).sum(dim=-1)
    var_ref = (probs_ref * torch.square(lprobs_score)).sum(dim=-1) - torch.square(mean_ref)
    discrepancy = (log_likelihood.sum(dim=-1) - mean_ref.sum(dim=-1)) / var_ref.sum(dim=-1).clamp_min(0.0001).sqrt()

    return discrepancy, log_likelihood.sum(dim=-1)


class ComputeStat(nn.Module):
    def __init__(self, model_name, dataset='xsum', device='cuda', cache_dir='./models', lora_r=None):
        super().__init__()
        self.device = device
        self.reference_model_name = get_model_fullname(model_name)
        self.scoring_model_name = get_model_fullname(model_name)

        def load_model(model_name, device, cache_dir, dtype_override=None):
            model_fullname = get_model_fullname(model_name)
            print(f'Loading model {model_fullname}...')
            model_kwargs = {}
            if model_name in float16_models:
                model_kwargs.update(dict(torch_dtype=torch.float16))
            # Gemma-1b's ~256k vocab makes fp32 logits/activations very large;
            # bf16 ~halves base-model + activation memory at no runtime cost.
            if 'gemma-1b' in model_name:
                model_kwargs.update(dict(torch_dtype=torch.bfloat16))
            # Explicit override (e.g. keep the trainable scoring model in fp32).
            if dtype_override is not None:
                model_kwargs.update(dict(torch_dtype=dtype_override))
            if torch.__version__ >= '2.0.0' and 'gemma' in model_name:
                model_kwargs.update({'attn_implementation': 'sdpa'})
            model = from_pretrained(AutoModelForCausalLM, model_fullname, model_kwargs, cache_dir, device=device)
            print(f'Moving model to {device}...', end='', flush=True)
            start = time.time()
            model.to(device)
            print(f'DONE ({time.time() - start:.2f}s)')
            return model

        # load scoring model (the trainable one). Keep gemma-1b in fp32 here so
        # the LoRA witness updates stay precise; the frozen reference below is bf16.
        self.scoring_tokenizer = load_tokenizer(model_name, dataset, cache_dir)
        scoring_dtype = GEMMA1B_SCORING_DTYPE if 'gemma-1b' in model_name else None
        scoring_model = load_model(model_name, device, cache_dir, dtype_override=scoring_dtype)
        if model_name in ['gemma-1b']:
            default_r, alpha, dropout = 4, 16, 0.05
        else:
            default_r, alpha, dropout = 8, 32, 0.1
        self.peft_config = LoraConfig(
            task_type=TaskType.CAUSAL_LM,
            inference_mode=False,
            r=lora_r if lora_r is not None else default_r,
            lora_alpha=alpha,
            lora_dropout=dropout,
            target_modules=["q_proj", "k_proj", "v_proj", "o_proj"],
        )
        self.scoring_model = get_peft_model(scoring_model, self.peft_config)

        # load sampling model
        self.reference_tokenizer = load_tokenizer(model_name, dataset, cache_dir)
        reference_model = load_model(model_name, device, cache_dir)
        self.reference_model = reference_model
        self.reference_model.eval()
        for p in self.reference_model.parameters():
            p.requires_grad = False

        total = sum(p.numel() for p in self.scoring_model.parameters())
        trainable = sum(p.numel() for p in self.scoring_model.parameters() if p.requires_grad)
        print(f"Trainable / total (parameters): {trainable}/{total}={trainable/total}")

        # Optional learned-domain estimator. When present, the special domain
        # names "estimate" / "softest" route a text through this classifier to
        # pick (or blend) the null distribution; see `_resolve_domain` /
        # `_predict_domain_probs` and `compute_p_value` below.
        self.domain_estimator = None

    def set_criterion_fn(self, criterion_fn):
        if criterion_fn == "mean":
            self.criterion = 'mean'
            self.criterion_fn = get_sampling_discrepancy_analytic
        else:
            raise ValueError(f"Unknown criterion function: {criterion_fn}")

    def print_gradient_requirement(self):
        for name, param in self.named_parameters():
            gradient_requirement = 'Requires Grad' if param.requires_grad else 'Does not require grad'
            color_code = '\033[92m' if param.requires_grad else '\033[91m'  # Green for requires grad, red for does not require grad
            reset_color = '\033[0m'  # Reset color after printing
            print(f"{name}: {color_code}{gradient_requirement}{reset_color}")

    def register_no_grad(self, module_names):
        for name, param in self.named_parameters():
            for selected_module in module_names:
                if selected_module in name:
                    param.requires_grad = False

    def save_pretrained(self, save_directory: str, save_null_distr_only=False):
        """
        Save the scoring model (with LoRA adapter) and all null_distr buffers in Hugging Face format.
        """
        os.makedirs(save_directory, exist_ok=True)

        # 1. Save the scoring model (LoRA adapter + base model).
        if not save_null_distr_only:
            scoring_dir = os.path.join(save_directory, "scoring_model")
            self.scoring_model.save_pretrained(scoring_dir, safe_serialization=True)

        # 2. Save every null_distr_* buffer.
        null_distrs = {}
        for buffer_name, buffer_value in self.named_buffers():
            if buffer_name.startswith("null_distr_"):
                domain = buffer_name.replace("null_distr_", "")
                null_distrs[domain] = buffer_value.detach().cpu()

        if null_distrs:
            torch.save(null_distrs, os.path.join(save_directory, "null_distrs.pt"))
            print(f"✅ Saved {len(null_distrs)} null distributions: {list(null_distrs.keys())}")

        # 3. Save config (including the domain list).
        config = {
            "domains": list(null_distrs.keys()),
            "criterion": getattr(self, "criterion", None),
        }
        with open(os.path.join(save_directory, "config.json"), "w") as f:
            json.dump(config, f)

        # 4. Save the domain estimator, if one has been trained.
        if not save_null_distr_only and self.domain_estimator is not None:
            self.domain_estimator.save_pretrained(save_directory)

        print(f"✅ Model saved to {save_directory}")

    @classmethod
    def from_pretrained(cls, load_directory: str, *args, **kwargs):
        """
        Load the scoring model, reference model, all null_distr buffers, and
        (if present) the domain estimator.
        """
        # 1. Construct the class.
        model = cls(*args, **kwargs)

        # 2. Load the scoring model.
        # NOTE: pass cache_dir through so the PEFT base model (google/gemma-3-1b-pt,
        # referenced in adapter_config.json) is resolved from the same cache used by
        # the constructor above, instead of silently re-downloading into the default
        # HF cache (~/.cache/huggingface). Pin to a single device instead of
        # device_map='auto' so this adapter checkpoint can't end up split across
        # devices from the reference model it sits next to.
        scoring_dir = os.path.join(load_directory, "scoring_model")
        model.scoring_model = AutoPeftModelForCausalLM.from_pretrained(
            scoring_dir,
            device_map={"": model.device},
            low_cpu_mem_usage=True,
            use_safetensors=True,
            cache_dir=kwargs.get("cache_dir"),
            trust_remote_code=True,
        )

        # 3. Load every null_distr.
        null_distrs_path = os.path.join(load_directory, "null_distrs.pt")
        if os.path.exists(null_distrs_path):
            null_distrs = torch.load(null_distrs_path, map_location="cpu")
            for domain, null_distr in null_distrs.items():
                model.set_null_distr(null_distr, domain)
            print(f"✅ Restored {len(null_distrs)} null distributions: {list(null_distrs.keys())}")

        # 4. Load config.
        config_path = os.path.join(load_directory, "config.json")
        if os.path.exists(config_path):
            with open(config_path, "r") as f:
                config = json.load(f)
            if "criterion" in config and config["criterion"] is not None:
                model.criterion = config["criterion"]
            print(f"✅ Loaded config: {config}")

        # 5. Load the domain estimator, if the checkpoint has one.
        # Pass our already-loaded `scoring_tokenizer` through: the classifier
        # shares the exact same base-model tokenizer (gemma-1b), so this
        # skips a second, redundant tokenizer load from disk/hub.
        #
        # This is wrapped in its own try/except: the domain estimator is an
        # *optional* feature (only needed for domain="estimate"/"softest").
        # A failure here -- e.g. an installed `transformers` version that
        # doesn't yet map the base model's config to a SequenceClassification
        # head -- must not take down the whole app; `model` (with manual
        # domain selection) should still load and serve requests.
        if os.path.exists(os.path.join(load_directory, DOMAIN_CLF_SUBDIR)):
            try:
                model.domain_estimator = DomainClassifier.from_pretrained(
                    load_directory,
                    cache_dir=kwargs.get("cache_dir", "./models"),
                    device=model.device,
                    tokenizer=model.scoring_tokenizer,
                )
            except Exception as e:  # noqa: BLE001 — degrade gracefully, see comment above
                print(
                    f"⚠️ Could not load domain estimator from {load_directory}: {e}\n"
                    f"   Continuing without it -- domain='estimate'/'softest' will be "
                    f"unavailable, but manual domain selection still works."
                )
                model.domain_estimator = None

        # Default to eval mode: a reloaded checkpoint is normally used for
        # inference, and leaving LoRA dropout active would perturb the witness
        # statistic away from the calibrated null distribution. Call .train()
        # explicitly before any further fine-tuning.
        model.scoring_model.eval()

        print(f"✅ Model loaded from {load_directory}")
        return model

    def compute_stats(self, tokenized=None, labels=[""], training_module=False):
        if training_module:
            logits_score = self.scoring_model(tokenized.input_ids, attention_mask=tokenized.attention_mask).logits[:,:-1,:]
            logits_ref = self.reference_model(tokenized.input_ids, attention_mask=tokenized.attention_mask).logits[:,:-1,:]
            crit, SPO_input  = self.criterion_fn(logits_ref, logits_score, labels)
        else:
            with torch.no_grad(): # get reference
                logits_score = self.scoring_model(tokenized.input_ids, attention_mask=tokenized.attention_mask).logits[:,:-1,:] # shape: [bsz, sentence_len, dim]
                logits_ref = self.reference_model(tokenized.input_ids, attention_mask=tokenized.attention_mask).logits[:,:-1,:]
                crit, SPO_input = self.criterion_fn(logits_ref, logits_score, labels)
        return crit, SPO_input, logits_score

    def forward(self, text, training_module=True):
        original_text = text[0]
        sampled_text = text[1]

        tokenized = self.scoring_tokenizer(original_text, return_tensors="pt", padding=True, return_token_type_ids=False).to(self.device)
        labels = tokenized.input_ids[:, 1:]
        train_original_crit, _, _ = self.compute_stats(tokenized, labels, training_module=training_module)

        tokenized = self.scoring_tokenizer(sampled_text, return_tensors="pt", padding=True, return_token_type_ids=False).to(self.device)
        labels = tokenized.input_ids[:, 1:]
        train_sampled_crit, _, _ = self.compute_stats(tokenized, labels, training_module=training_module)

        MMDloss = calculate_MMD_loss(train_original_crit, train_sampled_crit)
        output = dict(crit=[train_original_crit.detach(), train_original_crit, train_sampled_crit.detach(), train_sampled_crit], loss=MMDloss)
        return output

    def set_null_distr(self, null_distr: torch.Tensor, domain: str):
        """
        Set the null distribution tensor safely.
        """
        distr_name = f"null_distr_{domain}"
        self.register_buffer(distr_name, torch.empty(0))

        if not isinstance(null_distr, torch.Tensor):
            null_distr = torch.tensor(null_distr)

        # detach + clone + move to the right device
        null_distr = null_distr.detach().clone().to(self.device)

        # overwrite the buffer directly, to avoid issues with delattr
        self._buffers[distr_name] = null_distr
        print(f"✅ Null distribution on {domain} with shape: {self._buffers[distr_name].shape} with mean {self._buffers[distr_name].mean():.4f} and std {self._buffers[distr_name].std():.4f}")

    def _resolve_domain(self, text, domain: str):
        """Map the requested domain to a concrete one. The special name
        ``"estimate"`` predicts a single, most-likely domain from ``text``
        with the learned estimator (hard routing); any other name is returned
        unchanged. For the soft/mixture routing used by ``"softest"``, see
        ``_predict_domain_probs`` and ``compute_p_value_softest`` below."""
        if domain in (ESTIMATE_DOMAIN, "estimated"):
            if self.domain_estimator is None:
                raise ValueError(
                    "domain='estimate' requested but no domain estimator is "
                    "loaded. Train one with scripts/train_domain_clf.py (it is "
                    "saved into the checkpoint), or pass an explicit domain."
                )
            texts = [text] if isinstance(text, str) else list(text)
            return self.domain_estimator.predict(texts)[0]
        return domain

    def _predict_domain_probs(self, text):
        """Return ``{domain: probability}`` for ``text``, from the learned
        domain estimator, restricted to and renormalised over the domains
        that have a calibrated null distribution.

        This is the *soft* counterpart of ``_resolve_domain``'s ``"estimate"``:
        instead of collapsing the estimator's output to a single most-likely
        domain, it keeps the full predicted distribution, so a text that is
        itself a mixture of domains can be scored against a weighted
        combination of the candidate null distributions instead of being
        forced into exactly one of them.
        """
        if self.domain_estimator is None:
            raise ValueError(
                "domain='softest' requested but no domain estimator is "
                "loaded. Train one with scripts/train_domain_clf.py (it is "
                "saved into the checkpoint), or pass an explicit domain."
            )
        texts = [text] if isinstance(text, str) else list(text)
        probs = self.domain_estimator.predict_proba(texts)[0]

        available = set(self._null_distr_domains())
        probs = {d: p for d, p in probs.items() if d in available}
        z = sum(probs.values())
        if z <= 0:
            raise ValueError(
                "No overlap between the domain estimator's labels and the "
                "calibrated null distributions; cannot compute a 'softest' "
                "p-value."
            )
        return {d: p / z for d, p in probs.items()}

    def _compute_crit(self, text):
        """Tokenize ``text`` and compute the AdaJASA witness statistic. Shared
        by both the hard-domain (``compute_p_value``) and the soft/mixture
        (``compute_p_value_softest``) p-value paths."""
        tokenized = self.scoring_tokenizer(
            text,
            return_tensors="pt",
            padding=True,
            return_token_type_ids=False
        ).to(self.device)
        labels = tokenized.input_ids[:, 1:]

        with torch.inference_mode():
            crit, _, _ = self.compute_stats(tokenized, labels, training_module=False)
        return crit

    def compute_p_value(self, text, domain: str):
        """
        Compute p-value for given text using the null distribution of specified domain.

        Args:
            text: Input text to compute score for
            domain: Domain name to use for null distribution. Pass "estimate" to
                    let the learned estimator predict a single domain from the
                    text (hard routing), or "softest" to calibrate against a
                    weighted mixture of all calibrated domains' null
                    distributions -- weighted by the estimator's predicted
                    probabilities.
        """
        if domain == SOFTEST_DOMAIN:
            return self.compute_p_value_softest(text)

        domain = self._resolve_domain(text, domain)
        crit = self._compute_crit(text)

        # Look up the null distribution for this domain.
        distr_name = f"null_distr_{domain}"
        if not hasattr(self, distr_name):
            raise ValueError(
                f"No null distribution found for domain '{domain}'. "
                f"Available domains: {self.get_available_domains()}"
            )
        null_distr = getattr(self, distr_name)
        p_value = self.empirical_p_value(crit, null_distr)

        return crit, p_value

    def compute_p_value_softest(self, text):
        """Soft/mixture p-value:

            p-value = (1 + sum_k p_k * count_k) / (1 + sum_k p_k * m_k)

        where, for each calibrated domain k, ``p_k`` is the estimated
        probability that ``text`` belongs to domain k (from
        ``_predict_domain_probs``), ``m_k`` is the number of human-written
        calibration texts collected for domain k, and ``count_k`` is the
        number of those texts whose statistic falls below the observed
        statistic S(text).
        """
        domain_weights = self._predict_domain_probs(text)
        crit = self._compute_crit(text)

        numerator = 1.0
        denominator = 1.0
        for domain, weight in domain_weights.items():
            if weight <= 0:
                continue
            null_distr = getattr(self, f"null_distr_{domain}")
            m_k = null_distr.numel()
            count_k = (m_k - torch.searchsorted(null_distr, crit, right=False)[0]).item()
            numerator += weight * count_k
            denominator += weight * m_k

        p_value = torch.tensor(numerator / denominator, device=crit.device)
        return crit, p_value

    def empirical_p_value(self, crit: torch.Tensor, null_distr: torch.Tensor):
        # Compute p-value: (count + 1) / (total + 1)
        total = null_distr.numel()
        count = total - torch.searchsorted(null_distr, crit, right=False)[0]
        p_value = (count + 1.0) / (total + 1.0)
        return p_value

    def _null_distr_domains(self):
        """Domains with a concretely calibrated null distribution (excludes
        the pseudo-domain names ``"estimate"`` / ``"softest"``, which are
        resolved to a concrete domain -- or a weighted mixture of them -- at
        inference time)."""
        return [
            buffer_name.replace("null_distr_", "")
            for buffer_name in self._buffers.keys()
            if buffer_name.startswith("null_distr_")
        ]

    def get_available_domains(self):
        """
        Get list of all available domains with null distributions, plus the
        pseudo-domain names ("estimate", "softest") when a domain estimator
        is loaded.
        """
        domains = self._null_distr_domains()
        if getattr(self, "domain_estimator", None) is not None:
            domains.append(ESTIMATE_DOMAIN)
            domains.append(SOFTEST_DOMAIN)
        return domains


# Sub-directory (inside the AdaJASA checkpoint directory) that holds the learned
# domain estimator, so the classifier travels with the null distributions it
# selects between.
DOMAIN_CLF_SUBDIR = "domain_clf"

# Special domain name: route a text through the learned estimator instead of
# assuming the domain is known (oracle). Hard routing -- picks a single,
# most-likely domain (argmax).
ESTIMATE_DOMAIN = "estimate"

# Special domain name: soft/mixture routing. Instead of picking a single
# domain, calibrates against a weighted combination of every calibrated
# domain's null distribution, weighted by the domain estimator's predicted
# probabilities -- appropriate when the text may itself be a mixture of
# domains. See `ComputeStat.compute_p_value_softest`.
SOFTEST_DOMAIN = "softest"


class DomainClassifier(nn.Module):
    """A LoRA sequence classifier over text *domains*, sharing the gemma-1b base.

    Instead of assuming the test domain is known a priori, we estimate it from
    the text and let the predicted domain pick (or blend) which pre-computed
    AdaJASA null distribution is used for the decision.

    It is trained on ``(text, domain_label)`` pairs. Crucially, it never sees
    the human/machine label and is not part of null-distribution calibration,
    so it introduces no adaptivity into the p-values -- it only routes a test
    text to the correct calibration. A separate LoRA adapter on the *same*
    base model keeps this memory-efficient.
    """

    def __init__(
        self,
        model_name,
        label_names,
        device="cuda",
        cache_dir="./models",
        lora_r=8,
        max_length=512,
        tokenizer=None,
        _build_base=True,
    ):
        super().__init__()
        self.device = device
        self.model_name = model_name
        self.label_names = list(label_names)
        self.label2id = {name: i for i, name in enumerate(self.label_names)}
        self.id2label = {i: name for i, name in enumerate(self.label_names)}
        self.max_length = max_length
        # The classifier shares its base model with ComputeStat's
        # scoring/reference models (same `model_name`), so it can reuse an
        # already-loaded tokenizer (e.g. `ComputeStat.scoring_tokenizer`)
        # instead of loading -- or saving/reloading -- its own copy.
        self.tokenizer = tokenizer
        self.model = None

        # When loading via the classmethod `from_pretrained` below, the PEFT
        # model is restored from disk, so skip building a fresh base here
        # (avoids re-downloading / re-initialising the backbone).
        if not _build_base:
            return

        model_fullname = get_model_fullname(model_name)
        if self.tokenizer is None:
            tok_kwargs = {"padding_side": "right"}
            self.tokenizer = from_pretrained(AutoTokenizer, model_fullname, tok_kwargs, cache_dir=cache_dir)
        if self.tokenizer.pad_token_id is None:
            self.tokenizer.pad_token_id = self.tokenizer.eos_token_id

        base_kwargs = {
            "num_labels": len(self.label_names),
            "id2label": self.id2label,
            "label2id": self.label2id,
        }
        if "gemma-1b" in model_name:
            base_kwargs["torch_dtype"] = torch.bfloat16
        base_model = from_pretrained(AutoModelForSequenceClassification, model_fullname, base_kwargs, cache_dir, device=device)
        base_model.config.pad_token_id = self.tokenizer.pad_token_id

        peft_config = LoraConfig(
            task_type=TaskType.SEQ_CLS,
            inference_mode=False,
            r=lora_r,
            lora_alpha=lora_r * 4,
            lora_dropout=0.05,
            target_modules=["q_proj", "k_proj", "v_proj", "o_proj"],
        )
        self.model = get_peft_model(base_model, peft_config)
        self.model.to(device)

        trainable = sum(p.numel() for p in self.model.parameters() if p.requires_grad)
        total = sum(p.numel() for p in self.model.parameters())
        print(f"[DomainClassifier] {len(self.label_names)} domains; "
              f"trainable/total params: {trainable}/{total}={trainable / total:.4f}")

    def fit(self, texts, labels, epochs=3, lr=1e-4, batch_size=8, seed=42):
        """Train on ``(texts, labels)`` where each label is a domain name."""
        random.seed(seed)
        self.model.train()
        label_ids = [self.label2id[l] for l in labels]
        optimizer = AdamW(self.model.parameters(), lr=lr)
        n = len(texts)
        order = list(range(n))
        for epoch in range(epochs):
            random.shuffle(order)
            total_loss, correct, seen = 0.0, 0, 0
            for start in range(0, n, batch_size):
                idx = order[start:start + batch_size]
                batch_texts = [texts[i] for i in idx]
                batch_labels = torch.tensor([label_ids[i] for i in idx], device=self.device)
                enc = self.tokenizer(
                    batch_texts, return_tensors="pt", padding=True, truncation=True,
                    max_length=self.max_length, return_token_type_ids=False,
                ).to(self.device)
                optimizer.zero_grad()
                out = self.model(**enc, labels=batch_labels)
                out.loss.backward()
                optimizer.step()
                total_loss += out.loss.item() * len(idx)
                correct += (out.logits.argmax(dim=-1) == batch_labels).sum().item()
                seen += len(idx)
                if (start // batch_size) % 50 == 0:
                    torch.cuda.empty_cache()
            print(f"[DomainClassifier] epoch {epoch}: "
                  f"loss={total_loss / max(seen, 1):.4f} acc={correct / max(seen, 1):.4f}")
        return self

    @torch.no_grad()
    def predict(self, texts, batch_size=16):
        """Return a list of predicted domain *names* for ``texts``."""
        self.model.eval()
        preds = []
        for start in range(0, len(texts), batch_size):
            batch_texts = texts[start:start + batch_size]
            enc = self.tokenizer(
                batch_texts, return_tensors="pt", padding=True, truncation=True,
                max_length=self.max_length, return_token_type_ids=False,
            ).to(self.device)
            logits = self.model(**enc).logits
            preds.extend(self.id2label[i] for i in logits.argmax(dim=-1).tolist())
        return preds

    @torch.no_grad()
    def predict_proba(self, texts, batch_size=16):
        """Return a list of ``{domain_name: probability}`` dicts (softmax over
        the classifier's logits), one per text in ``texts``.

        Unlike :meth:`predict` (hard argmax), this keeps the full predicted
        distribution over domains -- what the soft/mixture p-value
        (domain="softest") needs: a text that is itself a blend of domains
        gets calibrated against a weighted combination of null distributions
        rather than forced into a single one.
        """
        self.model.eval()
        all_probs = []
        for start in range(0, len(texts), batch_size):
            batch_texts = texts[start:start + batch_size]
            enc = self.tokenizer(
                batch_texts, return_tensors="pt", padding=True, truncation=True,
                max_length=self.max_length, return_token_type_ids=False,
            ).to(self.device)
            logits = self.model(**enc).logits
            probs = torch.softmax(logits.float(), dim=-1)
            for row in probs.tolist():
                all_probs.append({self.id2label[i]: p for i, p in enumerate(row)})
        return all_probs

    def save_pretrained(self, ckpt_dir):
        """Save the LoRA adapter and label map under ``ckpt_dir``.

        The tokenizer is deliberately *not* saved here. It's byte-identical
        to the one `ComputeStat` already loads for its scoring/reference
        models (same base model), so writing another copy into
        ``domain_clf/`` would just duplicate ~40MB of vocab files per
        checkpoint for no benefit -- `from_pretrained` below reloads it from
        the shared ``cache_dir`` (or reuses a passed-in tokenizer) instead.
        """
        save_dir = os.path.join(ckpt_dir, DOMAIN_CLF_SUBDIR)
        os.makedirs(save_dir, exist_ok=True)
        self.model.save_pretrained(save_dir, safe_serialization=True)
        with open(os.path.join(save_dir, "label_names.json"), "w") as f:
            json.dump({"label_names": self.label_names, "base_model": self.model_name}, f)
        print(f"✅ Domain classifier saved to {save_dir} (labels: {self.label_names})")

    @classmethod
    def from_pretrained(cls, ckpt_dir, base_model=None, cache_dir="./models", device="cuda", max_length=512, tokenizer=None):
        """Load a saved domain classifier from ``ckpt_dir``. ``base_model`` is
        optional -- it's read from the checkpoint's ``label_names.json`` if
        omitted.

        Pass ``tokenizer=`` to reuse an already-loaded tokenizer (e.g.
        ``ComputeStat.scoring_tokenizer``) instead of loading a fresh copy --
        see the note on `save_pretrained` for why no tokenizer is bundled
        with this checkpoint in the first place.
        """
        save_dir = os.path.join(ckpt_dir, DOMAIN_CLF_SUBDIR)
        with open(os.path.join(save_dir, "label_names.json")) as f:
            meta = json.load(f)
        label_names = meta["label_names"]
        base_model = base_model or meta.get("base_model", "gemma-1b")

        obj = cls(
            base_model, label_names, device=device, cache_dir=cache_dir,
            max_length=max_length, tokenizer=tokenizer, _build_base=False,
        )
        if obj.tokenizer is None:
            model_fullname = get_model_fullname(base_model)
            tok_kwargs = {"padding_side": "right"}
            obj.tokenizer = from_pretrained(AutoTokenizer, model_fullname, tok_kwargs, cache_dir=cache_dir)
        if obj.tokenizer.pad_token_id is None:
            obj.tokenizer.pad_token_id = obj.tokenizer.eos_token_id

        # Match the training dtype (gemma-1b is trained/saved in bf16). Loaded
        # without device_map="auto" (we go straight to `.to(device, dtype)`
        # below) so there's no risk of the offload-split issue that motivates
        # the device pinning elsewhere in this file.
        dtype = torch.bfloat16 if "gemma-1b" in base_model else torch.float32
        obj.model = AutoPeftModelForSequenceClassification.from_pretrained(
            save_dir,
            num_labels=len(label_names),
            torch_dtype=dtype,
            low_cpu_mem_usage=True,
            cache_dir=cache_dir,
        )
        obj.model.to(device=device, dtype=dtype)
        obj.model.config.pad_token_id = obj.tokenizer.pad_token_id
        obj.model.eval()
        print(f"✅ Domain classifier loaded from {save_dir} (labels: {label_names}, dtype: {dtype})")
        return obj


def train_domain_clf(
    texts,
    labels,
    ckpt_dir,
    base_model="gemma-1b",
    cache_dir="./models",
    device="cuda",
    epochs=3,
    lr=1e-4,
    batch_size=8,
    lora_r=8,
    seed=42,
    tokenizer=None,
):
    """Train a LoRA domain estimator on ``(text, domain-label)`` pairs and save
    its checkpoint into ``ckpt_dir`` -- the *same* directory that holds the
    AdaJASA null distributions.

    Args:
        texts:     list[str] of input texts.
        labels:    list[str] of domain names, aligned with ``texts``.
        ckpt_dir:  AdaJASA checkpoint directory; the classifier is written to its
                   ``domain_clf/`` sub-directory.
        tokenizer: optional, already-loaded tokenizer to reuse (e.g. a
                   `ComputeStat` instance's `scoring_tokenizer`) instead of
                   loading a fresh copy of the same base-model tokenizer.

    Returns:
        The trained :class:`DomainClassifier`.
    """
    label_names = sorted(set(labels))
    clf = DomainClassifier(
        base_model, label_names, device=device, cache_dir=cache_dir, lora_r=lora_r,
        tokenizer=tokenizer,
    )
    clf.fit(texts, labels, epochs=epochs, lr=lr, batch_size=batch_size, seed=seed)
    clf.save_pretrained(ckpt_dir)
    return clf