File size: 38,266 Bytes
7c01d40
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830
831
832
833
834
835
836
837
838
839
840
841
842
843
844
845
846
847
848
849
850
851
852
853
854
855
856
857
858
859
860
861
862
863
864
865
866
867
868
869
870
871
872
873
874
875
876
877
878
879
880
881
882
883
884
885
886
887
888
889
890
891
892
893
894
#!/usr/bin/env python
"""IOL-AI 2026 submission -- International Linguistics Olympiad solver
  * Item-count detection now also understands lowercase sub-item letters and
    fill-in-the-blank underscores, which the previous regex set missed.
  * Answer voting no longer requires byte-identical normalized strings to
    agree -- it clusters samples by chrF similarity, so "iepure" and
    "iepure " and "Iepure" (or a stray trailing gloss) still count as the
    same vote instead of splitting support three ways.
  * Self-consistency sampling is convergence-aware: once a problem block's
    items all have a confident majority, that block is dropped from the
    prompt list for later rounds, so remaining time budget concentrates on
    the rows that are still undecided rather than re-rolling rows that have
    already settled.
  * The matching-task assignment solver is kept (free-form generation on
    permutation tasks tends to just emit the identity ordering) but is now
    reused for the "which digit" family of tasks too, with a lighter
    single-token constraint.
  * A rough token-length probe runs right after the tokenizer loads, so the
    starting batch size is chosen from the actual prompt lengths for this
    file instead of always starting at a fixed guess and paying for at least
    one guaranteed OOM.
Environment: transformers 4.44.1 / torch 2.4.0 / autoawq, 16GB T4, fp16 only,
no internet, hard 30-minute wall clock.
"""
import os
import re
import json
import time
import unicodedata
from collections import Counter, defaultdict

START = time.time()

# ---------------------------------------------------------------------------
# Budget
# ---------------------------------------------------------------------------
# The harness kills the process at the wall-clock limit with zero credit for
# work in flight, so every stage aims to finish with a margin, not exactly on
# time. SAFETY is subtracted up front rather than checked ad hoc so every
# caller of time_left() sees the same conservative deadline.
TIME_LIMIT = float(os.environ.get("IOL_TIME_LIMIT", "1800"))
SAFETY = float(os.environ.get("IOL_SAFETY", "150"))
DEADLINE = START + TIME_LIMIT - SAFETY

TEST_CSV = os.environ.get("IOL_TEST_CSV", "/tmp/data/test.csv")
OUT_CSV = os.environ.get("IOL_OUT_CSV", "submission.csv")
MODEL_ID = os.environ.get("IOL_MODEL", ".")
WANT_EXPLANATION = os.environ.get("IOL_EXPLAIN", "1") == "1"
MAX_NEW = int(os.environ.get("IOL_MAXNEW", "900"))
MAX_SAMPLES = int(os.environ.get("IOL_MAXSAMPLES", "8"))

# Earlier runs found the organizers' plain, no-chain-of-thought prompt style
# outscored an elaborate CoT harness on the hidden set. Default to that
# "faithful" style, but keep the richer CoT prompt available for A/B use --
# the parsing/voting/repair layers below are written to help either mode
# rather than assuming one.
PROMPT_STYLE = os.environ.get("IOL_PROMPT_STYLE", "faithful")  # faithful | guided
SAMPLE_TEMP = float(os.environ.get("IOL_TEMP", "0.5"))

os.environ.setdefault("HF_HUB_OFFLINE", "1")
os.environ.setdefault("TRANSFORMERS_OFFLINE", "1")
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")


def log(msg):
    print(f"[{time.time() - START:7.1f}s] {msg}", flush=True)


def time_left():
    return DEADLINE - time.time()


# ---------------------------------------------------------------------------
# How many items does this problem block ask for?
# ---------------------------------------------------------------------------
_NUM_LINE = re.compile(r"^[ \t]*(\d{1,3})[.)\]]", re.M)
_NUM_PAREN = re.compile(r"\((\d{1,3})\)")
_NUM_RANGE = re.compile(r"\(?(\d{1,3})\s*(?:[-\u2013\u2014]|to)\s*(\d{1,3})\)?")
_LETTER_LINE = re.compile(r"^[ \t]*([A-Za-z])[.)\]]\s", re.M)
_LETTER_PAREN = re.compile(r"\(([A-Za-z])\)")
_BLANK_MARK = re.compile(r"_{2,}|\[\s*(?:blank|___*)\s*\]", re.I)


def count_items(query, task_type="", context=""):
    """Best-effort count of the sub-answers a block expects. Never returns 0."""
    q = query or ""

    nums = set(int(m) for m in _NUM_LINE.findall(q)) | set(
        int(m) for m in _NUM_PAREN.findall(q))
    letters = set(_LETTER_LINE.findall(q)) | set(_LETTER_PAREN.findall(q))

    span = 0
    for a, b in _NUM_RANGE.findall(q):
        a, b = int(a), int(b)
        if 0 < b - a < 60:
            span = max(span, b - a + 1)

    marker_count = max(len(nums), len(letters))
    if span and marker_count and span != marker_count:
        # Explicit markers on the page are what must actually be answered,
        # even if a stated range like "items 1-5" disagrees with them.
        n = marker_count
    else:
        n = max(span, marker_count)

    if n > 1:
        return n

    blanks = len(_BLANK_MARK.findall(q))
    if blanks > 1:
        return blanks

    # No numbering at all: "Translate the following into Kalam:" followed by
    # one bare item per line.
    body_lines = [l.strip() for l in q.splitlines() if l.strip()]
    if len(body_lines) > 1:
        head, rest = body_lines[0], body_lines[1:]
        if rest and (head.endswith((":", ".")) or len(rest) >= 2):
            return len(rest)

    # A bare instruction ("State the correspondence for each pair.") whose
    # items actually live in the shared context block.
    if context:
        cn = len(set(int(m) for m in _NUM_LINE.findall(context)))
        if cn > 1:
            return cn
        cl = len(set(_LETTER_LINE.findall(context)))
        if cl > 1:
            return cl

    return max(n, 1)


# ---------------------------------------------------------------------------
# Cleaning and parsing model output into exactly n answers
# ---------------------------------------------------------------------------
_LEAD_MARK = re.compile(r"^\s*(?:\(?\d{1,3}\)?[.):\]]\s*|[-*\u2022]\s+)")
_CODE_FENCE = re.compile(r"^```[a-zA-Z]*\s*$")
_META_TALK = re.compile(
    r"^\s*(?:here (?:are|is)\b|answers?\s*:?\s*$|explanation\b|note\b|okay\b|"
    r"solution\b|reasoning\b|analysis\b|translations?\s*:?\s*$|the answers?\b|"
    r"let me\b|first,|so,|therefore\b|thus\b|to summarize\b)",
    re.I,
)


def tidy(line):
    s = line.strip()
    s = _LEAD_MARK.sub("", s)
    s = s.strip().strip("`").strip()
    if len(s) >= 2 and s[0] == s[-1] and s[0] in "\"'\u201c\u201d":
        s = s[1:-1].strip()
    return s.strip()


def source_fallback(query, n):
    """Last-resort answers when the model produced nothing usable.
    A blank answer scores zero on both exact-match and chrF, but echoing the
    prompt's own item text is nearly free and, on transcription-adjacent
    tasks, shares enough surface form with the true answer to pick up
    partial chrF credit. Better than nothing, never claimed to be more.
    """
    q = query or ""
    picked = []
    for ln in q.splitlines():
        s = ln.strip()
        if not s:
            continue
        m = re.match(r"^\(?(\d{1,3})\)?[.):\]]\s*(.+)$", s)
        if m:
            picked.append(m.group(2).strip())
    if not picked:
        lines = [l.strip() for l in q.splitlines() if l.strip()]
        if len(lines) > 1 and lines[0].endswith((":", ".")):
            picked = lines[1:]
    picked = [p.split("|")[0].strip() if "|" in p else p for p in picked]
    picked = [p for p in picked if p]
    while len(picked) < n:
        picked.append(picked[-1] if picked else "?")
    return picked[:n]


def squeeze_to_n(items, n, fallback=None):
    items = [i for i in items if i and i.strip()]
    if len(items) > n:
        # The model reasons first and answers last; absent an explicit
        # marker to slice on, the tail of the output is the answer block.
        items = items[-n:]
    while len(items) < n:
        if fallback and len(items) < len(fallback):
            items.append(fallback[len(items)])
        else:
            items.append(items[-1] if items else "?")
    return items[:n]


def parse_model_output(text, n, fallback=None):
    if not text or not text.strip():
        return list(fallback[:n]) if fallback else ["?"] * n

    marker = None
    for m in re.finditer(r"(?:^|\n)\s*(?:final\s+)?answers?\s*:\s*\n?", text, re.I):
        marker = m
    body = text[marker.end():] if marker else text

    labelled, plain = [], []
    for ln in body.splitlines():
        if _CODE_FENCE.match(ln):
            continue
        m = re.match(r"^\s*\(?(\d{1,3})\)?[.):\]]\s*(.+)$", ln.strip())
        if m:
            val = tidy(m.group(2))
            if val and not _META_TALK.match(val):
                labelled.append((int(m.group(1)), val))
        c = tidy(ln)
        if c and not _META_TALK.match(c):
            plain.append(c)

    if len(labelled) >= n:
        # Numbered labels, when present, are a stronger placement signal than
        # line order -- keep the last value written under each label since
        # models sometimes restate an earlier answer while reasoning.
        by_label = {}
        for lab, val in labelled:
            by_label[lab] = val
        keys = sorted(by_label)
        if len(keys) >= n:
            return [by_label[k] for k in keys[:n]]

    return squeeze_to_n(plain, n, fallback)


_DIGITS_ONLY = re.compile(r"-?\d+")


def coerce_numeral(answer, task_type):
    """text_to_num answers should be bare digits; pull them out if the model
    wrapped them in words or punctuation instead of reworking the parser for
    one task family."""
    if (task_type or "").strip().lower() != "text_to_num":
        return answer
    m = _DIGITS_ONLY.search(answer or "")
    return m.group(0) if m else answer


def normalize(s):
    s = unicodedata.normalize("NFC", (s or "").strip().lower())
    s = re.sub(r"\s+", " ", s)
    return s.strip(" .!?;:,")


# ---------------------------------------------------------------------------
# chrF, computed inline since sacrebleu may not be present in the sandbox
# ---------------------------------------------------------------------------
def _char_ngrams(s, k):
    s = re.sub(r"\s+", "", s)
    if len(s) < k:
        return Counter()
    return Counter(s[i:i + k] for i in range(len(s) - k + 1))


def chrf(hyp, ref, order=6, beta=2.0):
    if not hyp or not ref:
        return 0.0
    precisions, recalls = [], []
    for k in range(1, order + 1):
        h, r = _char_ngrams(hyp, k), _char_ngrams(ref, k)
        if not h or not r:
            continue
        overlap = sum((h & r).values())
        precisions.append(overlap / max(1, sum(h.values())))
        recalls.append(overlap / max(1, sum(r.values())))
    if not precisions:
        return 0.0
    p = sum(precisions) / len(precisions)
    r = sum(recalls) / len(recalls)
    if p + r == 0:
        return 0.0
    b2 = beta * beta
    return (1 + b2) * p * r / (b2 * p + r)


# ---------------------------------------------------------------------------
# Voting: cluster candidates by similarity rather than exact string match
# ---------------------------------------------------------------------------
CLUSTER_THRESHOLD = 0.90  # chrF similarity above which two answers "agree"


def cluster_candidates(cands):
    """Group near-duplicate strings together (single-linkage on chrF).
    Exact-match voting undercounts agreement: "iepurele" and "iepurele " or a
    stray trailing gloss are the same underlying answer but would previously
    split into separate buckets and never reach a majority. Grouping by
    similarity instead means minor formatting noise doesn't cost a block its
    vote.
    """
    groups = []  # list of [members]
    for c in cands:
        placed = False
        for g in groups:
            if chrf(c, g[0]) >= CLUSTER_THRESHOLD or normalize(c) == normalize(g[0]):
                g.append(c)
                placed = True
                break
        if not placed:
            groups.append([c])
    return groups


def vote_answer(cands, anchor=None):
    """Choose one answer for an item from the greedy pass plus samples.
    The greedy (temperature 0) answer is the default and can only be
    displaced by a cluster that (a) has at least two members and (b) is
    strictly larger than the cluster the anchor itself belongs to. This
    asymmetry matters: with only a handful of samples, treating all
    candidates as equal and breaking ties by "closest to the others" is
    poorly conditioned -- with two candidates that comparison is symmetric
    and degenerates into picking whichever string happens to be shorter. On
    a held-out mock set that variant swapped out the greedy answer for a
    sampled one roughly half the time and cost a large chunk of exact match.
    Anchoring keeps voting monotone: it only overrides on genuine, repeated
    agreement.
    """
    cands = [c for c in cands if c and c.strip()]
    if anchor is None:
        anchor = cands[0] if cands else "?"
    if len(cands) < 3:
        return anchor

    groups = cluster_candidates(cands)
    anchor_group_size = 1
    for g in groups:
        if any(normalize(m) == normalize(anchor) or chrf(m, anchor) >= CLUSTER_THRESHOLD
               for m in g):
            anchor_group_size = len(g)
            break

    best_group = max(groups, key=len)
    if len(best_group) >= 2 and len(best_group) > anchor_group_size:
        # Within the winning cluster, prefer whichever exact form is most
        # common; ties fall back to the first (greedy-adjacent) occurrence.
        return Counter(best_group).most_common(1)[0][0]
    return anchor


# ---------------------------------------------------------------------------
# Matching / assignment problems: read next-token logits instead of trusting
# free-form generation, which tends to just emit the option letters in order.
# ---------------------------------------------------------------------------
_ROW_OPT = re.compile(r"^[ \t]*([A-Za-z])[.)]\s+(.+)$", re.M)
_ROW_ITEM = re.compile(r"^[ \t]*(\d{1,3})[.)]\s+(.+)$", re.M)


def split_matching_block(context):
    items = [(int(a), b.strip()) for a, b in _ROW_ITEM.findall(context or "")]
    opts = [(a, b.strip()) for a, b in _ROW_OPT.findall(context or "")]
    seen = set()
    items = [x for x in items if not (x[0] in seen or seen.add(x[0]))]
    seen = set()
    opts = [x for x in opts if not (x[0] in seen or seen.add(x[0]))]
    return items, opts


def assign_max_weight(score):
    """One-to-one assignment maximizing total score. scipy if available,
    otherwise a greedy pass plus a few rounds of local 2-swaps."""
    n, m = len(score), len(score[0])
    try:
        import numpy as _np
        from scipy.optimize import linear_sum_assignment
        _, cols = linear_sum_assignment(-_np.array(score))
        return list(cols)
    except Exception:
        pass
    used, out = set(), [0] * n
    order = sorted(
        range(n),
        key=lambda i: -(max(score[i]) - sorted(score[i])[-2] if m > 1 else 0),
    )
    for i in order:
        j = max((j for j in range(m) if j not in used), key=lambda j: score[i][j],
                default=0)
        used.add(j)
        out[i] = j
    for _ in range(4):
        improved = False
        for a in range(n):
            for b in range(a + 1, n):
                cur = score[a][out[a]] + score[b][out[b]]
                alt = score[a][out[b]] + score[b][out[a]]
                if alt > cur + 1e-9:
                    out[a], out[b] = out[b], out[a]
                    improved = True
        if not improved:
            break
    return out


def enforce_bijection(answers):
    """If every answer is a single letter and duplicates exist where the full
    alphabet of options wasn't used, reassign duplicated slots to the unused
    letters. Tightly scoped so it never touches anything but this exact
    single-letter-permutation shape."""
    if len(answers) < 3 or not all(re.fullmatch(r"[A-Za-z]", a or "") for a in answers):
        return answers
    n = len(answers)
    universe = [chr(ord("A") + i) for i in range(n)]
    upper = [a.upper() for a in answers]
    if len(set(upper)) == n:
        return answers
    unused = [l for l in universe if l not in set(upper)]
    if not unused:
        return answers
    seen, out = set(), []
    for a in upper:
        if a in seen and unused:
            out.append(unused.pop(0))
        else:
            seen.add(a)
            out.append(a)
    return out


# ---------------------------------------------------------------------------
# Prompting
# ---------------------------------------------------------------------------
FAITHFUL_SYSTEM = (
    "You solve International Linguistics Olympiad problems. Answer every "
    "numbered item. Put each answer on its own line, in order, with no "
    "numbering and no extra text."
)

GUIDED_SYSTEM = (
    "You are a gold medallist at the International Linguistics Olympiad.\n"
    "Each problem gives data from a language you have never seen before; "
    "everything needed is in the problem itself, and no outside knowledge is "
    "required.\n"
    "Work method: align the given examples, segment the forms, find the "
    "recurring morphemes and the rule that orders them, verify the rule "
    "against every example given, then apply it to the items you are asked "
    "for.\n"
    "Reason concisely. Finish with a block that starts on its own line with "
    "exactly ANSWERS: followed by one answer per line, in order, with no "
    "numbering, no commentary, and no blank lines.\n"
    "Always give your best guess for every item -- never leave one blank."
)

TASK_HINTS = {
    "translation": "Give only the translation itself -- no source text, no "
                   "gloss, no quotation marks.",
    "match_letters": "Give a single capital letter per item, identifying its "
                     "match. Every letter is used exactly once; none repeats.",
    "fill_blanks": "Give only the missing form for that blank -- not the "
                   "whole line, not a gloss.",
    "text_to_num": "Give the number in digits only (e.g. 111).",
    "num_to_text": "Give the number spelled out in words, in the problem "
                   "language.",
}


def build_prompt(row, n):
    if PROMPT_STYLE == "faithful":
        return f"{row['context'].strip()}\n\n{row['query'].strip()}"
    hint = TASK_HINTS.get((row.get("task_type") or "").strip().lower(), "")
    return (
        f"{row['context'].strip()}\n\n{row['query'].strip()}\n\n"
        f"There are exactly {n} item{'s' if n != 1 else ''} to answer."
        + (f" {hint}" if hint else "")
        + f"\nAfter reasoning, write ANSWERS: on its own line, then exactly "
          f"{n} line{'s' if n != 1 else ''}, one answer per item, in order."
    )


EXPLAIN_SYSTEM = (
    "You explain International Linguistics Olympiad solutions to a human "
    "judge. Given a problem and the answers produced, state the key rules of "
    "the language that justify them -- relevant morphemes, word order, any "
    "sound changes. Be specific and brief (2-4 sentences or a few short "
    "bullets); do not restate a chain of reasoning."
)


def build_explain_prompt(row, answers):
    return (
        f"{row['context'].strip()}\n\n{row['query'].strip()}\n\n"
        "Answers given:\n" + "\n".join(f"- {a}" for a in answers) +
        "\n\nBriefly explain the linguistic rules behind these answers."
    )


# ---------------------------------------------------------------------------
# Offline diagnostics (never runs on the platform -- gold labels aren't
# present there, so IOL_GOLD stays unset)
# ---------------------------------------------------------------------------
def report_dev_score(preds):
    gold_path = os.environ.get("IOL_GOLD")
    if not gold_path or not os.path.exists(gold_path):
        return
    try:
        import ast
        import pandas as pd
        gold_df = pd.read_csv(gold_path, dtype=str)
        exact, chrfs = [], []
        for _, r in gold_df.iterrows():
            gold = ast.literal_eval(r["answer"])
            p = preds.get(str(r["id"]), [])
            p = list(p)[:len(gold)] + [""] * max(0, len(gold) - len(p))
            for gi, pi in zip(gold, p):
                alts = gi if isinstance(gi, (list, tuple)) else [gi]
                alts = [str(a) for a in alts]
                exact.append(1.0 if any(pi.strip() == a.strip() for a in alts) else 0.0)
                chrfs.append(max(chrf(pi, a) for a in alts))
        em = sum(exact) / max(1, len(exact))
        cf = sum(chrfs) / max(1, len(chrfs))
        log(f"  [dev] EM={em:.4f} chrF~={cf:.4f} score~={(em * cf) ** 0.5:.4f} "
            f"over {len(exact)} items")
    except Exception as e:
        log(f"  [dev] scoring failed: {type(e).__name__}: {e}")


def write_submission(path, ids, preds, explanations=None):
    import pandas as pd
    rows = []
    for i in ids:
        rec = {"id": i, "pred": json.dumps(preds[i], ensure_ascii=False)}
        if explanations is not None:
            rec["explanation"] = explanations.get(i, "")
        rows.append(rec)
    pd.DataFrame(rows).to_csv(path, index=False)


# ---------------------------------------------------------------------------
# Main
# ---------------------------------------------------------------------------
def main():
    import pandas as pd

    df = pd.read_csv(TEST_CSV, dtype=str).fillna("")
    ids = [str(x) for x in df["id"].tolist()]
    ns = [count_items(r.get("query", ""), r.get("task_type", ""), r.get("context", ""))
          for _, r in df.iterrows()]
    total = sum(ns)
    log(f"loaded {len(df)} problems, {total} items "
        f"(min={min(ns)} max={max(ns)} mean={total / len(ns):.1f})")

    fallbacks = {i: source_fallback(r.get("query", ""), n)
                 for i, (_, r), n in zip(ids, df.iterrows(), ns)}

    # 1. A complete, correctly-shaped submission before the model even loads.
    preds = {i: list(fallbacks[i]) for i in ids}
    explanations = {i: "" for i in ids} if WANT_EXPLANATION else None
    write_submission(OUT_CSV, ids, preds, explanations)
    log(f"wrote placeholder {OUT_CSV} ({len(ids)} rows)")

    # 2. Load model.
    import torch
    from transformers import (AutoTokenizer, AutoModelForCausalLM,
                              StoppingCriteria, StoppingCriteriaList)

    class Deadline(StoppingCriteria):
        def __init__(self, stop_at):
            self.stop_at = stop_at

        def __call__(self, input_ids, scores, **kw):
            return time.time() > self.stop_at

    log("loading tokenizer/model ...")
    tok = AutoTokenizer.from_pretrained(MODEL_ID, trust_remote_code=True)
    if tok.pad_token is None:
        tok.pad_token = tok.eos_token
    tok.padding_side = "left"

    def _load(dev_map):
        try:
            return AutoModelForCausalLM.from_pretrained(
                MODEL_ID, torch_dtype=torch.float16, device_map=dev_map,
                trust_remote_code=True).eval()
        except TypeError:
            return AutoModelForCausalLM.from_pretrained(
                MODEL_ID, dtype=torch.float16, device_map=dev_map,
                trust_remote_code=True).eval()

    try:
        # Pin every layer to the GPU explicitly. device_map="auto" is free to
        # spill a layer or two to CPU when VRAM looks tight, which is nearly
        # invisible in logs but makes generation roughly two orders of
        # magnitude slower -- the worst possible failure mode under a hard
        # deadline. Only fall back to "auto" if the explicit placement fails
        # outright.
        model = _load({"": 0} if torch.cuda.is_available() else "auto")
    except Exception as e:
        log(f"pinned load failed ({type(e).__name__}: {e}); falling back to auto")
        model = _load("auto")

    devs = set(str(p.device) for p in model.parameters())
    log(f"model ready on {sorted(devs)} ({time_left():.0f}s of budget left)")
    if any(d.startswith("cpu") or d == "meta" for d in devs):
        log("WARNING: part of the model is off-GPU; generation will be very slow")
    if torch.cuda.is_available():
        log(f"  VRAM allocated {torch.cuda.memory_allocated()/1e9:.2f} GB / "
            f"{torch.cuda.get_device_properties(0).total_memory/1e9:.1f} GB")

    prompts = []
    for (_, r), n in zip(df.iterrows(), ns):
        system = FAITHFUL_SYSTEM if PROMPT_STYLE == "faithful" else GUIDED_SYSTEM
        msgs = [{"role": "system", "content": system},
                {"role": "user", "content": build_prompt(r, n)}]
        prompts.append(tok.apply_chat_template(msgs, tokenize=False,
                                               add_generation_prompt=True))

    # Probe actual prompt lengths so the starting batch size reflects this
    # file's problems instead of a fixed guess that either wastes headroom on
    # short prompts or guarantees an OOM on long ones.
    sample_lens = [len(tok.encode(p, add_special_tokens=False))
                   for p in prompts[: min(8, len(prompts))]]
    approx_max_len = max(sample_lens) if sample_lens else 512
    default_batch = 4 if approx_max_len < 2500 else (2 if approx_max_len < 4500 else 1)
    batch_size = int(os.environ.get("IOL_BATCH", str(default_batch)))
    log(f"prompt length probe: ~{approx_max_len} tok max of {len(sample_lens)} "
        f"sampled -> starting batch_size={batch_size}")

    def generate(texts, max_new, sample, temp=0.7):
        nonlocal batch_size
        out = [""] * len(texts)
        order = sorted(range(len(texts)), key=lambda i: len(texts[i]))
        i = 0
        while i < len(order):
            if time_left() < 25:
                log("  out of time inside generate(); returning partial batch")
                break
            idx = order[i:i + batch_size]
            chunk = [texts[j] for j in idx]
            try:
                enc = tok(chunk, return_tensors="pt", padding=True,
                          truncation=True, max_length=6144).to(model.device)
                # repetition_penalty=1.0 explicitly. Several AWQ chat models
                # ship a generation_config.json with a >1.0 repetition
                # penalty baked in, and unlike temperature/top_p (which
                # greedy decoding ignores, with a warning) a repetition
                # penalty is applied under greedy decoding silently. A large
                # share of gold answers in agglutinative languages legitimately
                # repeat a character run, so leaving the shipped penalty in
                # place pushes generation away from exactly the strings
                # needed.
                kw = dict(max_new_tokens=max_new, pad_token_id=tok.pad_token_id,
                          repetition_penalty=1.0,
                          stopping_criteria=StoppingCriteriaList(
                              [Deadline(DEADLINE - 10)]))
                if sample:
                    kw.update(do_sample=True, temperature=temp, top_p=0.95)
                else:
                    kw.update(do_sample=False)
                with torch.no_grad():
                    o = model.generate(**enc, **kw)
                for k, j in enumerate(idx):
                    out[j] = tok.decode(o[k][enc["input_ids"].shape[1]:],
                                        skip_special_tokens=True)
                i += batch_size
            except torch.cuda.OutOfMemoryError:
                torch.cuda.empty_cache()
                if batch_size == 1:
                    log("  OOM at batch=1; skipping this item")
                    i += 1
                else:
                    batch_size = max(1, batch_size // 2)
                    log(f"  OOM -> batch_size={batch_size}")
            except Exception as e:
                log(f"  generate error: {type(e).__name__}: {e}")
                i += batch_size
        return out

    def solve_by_assignment(row, n):
        """Score every (item, option-letter) pair via next-token logits and
        take the best one-to-one assignment, instead of trusting free-form
        text for match_letters problems (which tends to just emit A, B, C...
        in order -- a valid permutation, so no repair catches it, and it
        scores near zero).
        Raw next-token logprobs are known to carry a systematic per-token
        prior independent of content (Zhao et al. 2021, "Calibrate Before
        Use") -- e.g. a model may just prefer the token for "C" regardless of
        which option is actually correct. We estimate that per-letter prior
        with a content-free version of the same prompt and subtract it before
        assignment, so the score reflects evidence from the item text rather
        than a baked-in letter preference.
        """
        items, opts = split_matching_block(row.get("context", ""))
        if len(items) < 3 or len(opts) < 3 or len(items) != n:
            return None
        letters = [o[0] for o in opts]
        cand_ids = []
        for L in letters:
            token_ids = set()
            for form in (L, " " + L):
                enc = tok.encode(form, add_special_tokens=False)
                if enc:
                    token_ids.add(enc[0])
            cand_ids.append(sorted(token_ids))

        ctx = row["context"].strip()

        def _logprob_batch(msgs_list):
            """Run one forward pass per prompt in msgs_list, return a list of
            per-candidate-letter logprob lists (same shape as `scores` rows)."""
            texts = [tok.apply_chat_template(m, tokenize=False,
                                              add_generation_prompt=True)
                     for m in msgs_list]
            out = []
            step = 4
            for s0 in range(0, len(texts), step):
                if time_left() < 30:
                    return None
                chunk = texts[s0:s0 + step]
                enc = tok(chunk, return_tensors="pt", padding=True,
                          truncation=True, max_length=6144).to(model.device)
                with torch.no_grad():
                    logits = model(**enc).logits[:, -1, :].float()
                logp = torch.log_softmax(logits, dim=-1)
                for b in range(len(chunk)):
                    out.append([max(logp[b, i].item() for i in ids_)
                                for ids_ in cand_ids])
            return out

        # --- real per-item prompts -------------------------------------
        q_msgs = []
        for num, itext in items:
            q_msgs.append([
                {"role": "system", "content":
                 "You match items to their correct counterparts in a "
                 "linguistics problem. Reply with one option letter only."},
                {"role": "user", "content":
                 f"{ctx}\n\nWhich lettered option corresponds to item {num} "
                 f"({itext})? Reply with the option letter only."},
            ])
        scores = _logprob_batch(q_msgs)
        if scores is None:
            return None

        # --- content-free baseline prompt -------------------------------
        # Same instruction and option list (letters need to stay visible
        # since the token candidates depend on the model's tokenization in
        # context), but with the item text and its identifying number
        # blanked out, so the only thing driving the letter choice is the
        # model's own prior over "A"/"B"/"C"... in this response format --
        # exactly the confound we want to remove from `scores`.
        opts_block = "\n".join(f"{L}. {t}" for L, t in opts)
        baseline_msgs = [{"role": "system", "content":
                          "You match items to their correct counterparts in a "
                          "linguistics problem. Reply with one option letter only."},
                         {"role": "user", "content":
                          f"{opts_block}\n\nWhich lettered option corresponds "
                          f"to item ___ (___)? Reply with the option letter "
                          f"only."}]
        baseline_scores = _logprob_batch([baseline_msgs])
        if baseline_scores is None:
            # Baseline pass failed/timed out -- fall back to uncalibrated
            # scores rather than losing the whole solver over one extra call.
            log("  calibration baseline unavailable; using raw scores")
            calibrated = scores
        else:
            baseline = baseline_scores[0]  # one row, per-letter logprob
            calibrated = [[s - baseline[j] for j, s in enumerate(row_scores)]
                          for row_scores in scores]

        cols = assign_max_weight(calibrated)
        return [letters[c] for c in cols]

    def apply_parsers(row_id, row, n, text):
        answer = enforce_bijection(parse_model_output(text, n, fallbacks[row_id]))
        task_type = (row.get("task_type") or "").strip().lower()
        if task_type == "match_letters":
            try:
                mm = solve_by_assignment(row, n)
                if mm and len(mm) == n:
                    return mm
            except Exception as e:
                log(f"  assignment solver failed on {row_id}: {type(e).__name__}: {e}")
        if task_type == "text_to_num":
            answer = [coerce_numeral(a, task_type) for a in answer]
        return answer

    # 3. Pass 1: greedy decoding, guarantees a complete answer set.
    TOK_PER_S = float(os.environ.get("IOL_TOKS", "30"))
    adaptive = int(0.40 * max(1.0, time_left()) * TOK_PER_S / max(1, len(df)))
    max_new = max(192, min(MAX_NEW, adaptive))
    log(f"reasoning budget: {max_new} new tokens/problem "
        f"(adaptive={adaptive}, cap={MAX_NEW}, {len(df)} problems)")

    t0 = time.time()
    texts = generate(prompts, max_new=max_new, sample=False)
    pass1_cost = time.time() - t0

    samples = {i: [] for i in ids}
    match_rows = 0
    for (row_id, n, txt), (_, row) in zip(zip(ids, ns, texts), df.iterrows()):
        a = apply_parsers(row_id, row, n, txt)
        if (row.get("task_type") or "").strip().lower() == "match_letters":
            match_rows += 1
        preds[row_id] = a
        samples[row_id].append(a)
    if match_rows:
        log(f"assignment solver attempted on {match_rows} match_letters problem(s)")
    write_submission(OUT_CSV, ids, preds, explanations)
    no_block = sum(1 for t in texts if not re.search(r"answers?\s*:", t or "", re.I))
    empty = sum(1 for t in texts if not (t or "").strip())
    log(f"pass 1 (greedy) done in {pass1_cost:.0f}s -> submission written "
        f"({no_block}/{len(texts)} without an ANSWERS: block, {empty} empty)")
    report_dev_score(preds)

    # 4. Convergence-aware self-consistency: only re-sample rows whose items
    # don't yet have a confident majority, so budget concentrates on the rows
    # that are actually still undecided instead of re-rolling settled ones.
    def row_converged(row_id):
        s = samples[row_id]
        if len(s) < 3:
            return False
        for k in range(len(s[0])):
            col = [x[k] for x in s if k < len(x)]
            groups = cluster_candidates(col)
            if not any(len(g) >= 2 for g in groups):
                return False
        return True

    reserve = 0.0
    if WANT_EXPLANATION:
        reserve = min(300.0, 0.25 * pass1_cost + 60)

    active_ids = list(ids)
    n_extra = 0
    while time_left() - reserve > pass1_cost * 0.5 and n_extra < MAX_SAMPLES and active_ids:
        n_extra += 1
        active_ids = [i for i in active_ids if not row_converged(i)]
        if not active_ids:
            log("  every row has converged; stopping self-consistency early")
            break
        log(f"self-consistency pass {n_extra} on {len(active_ids)}/{len(ids)} "
            f"undecided rows ({time_left():.0f}s left)")
        pos = {rid: k for k, rid in enumerate(ids)}
        sub_prompts = [prompts[pos[i]] for i in active_ids]
        texts = generate(sub_prompts, max_new=max_new, sample=True, temp=SAMPLE_TEMP)
        row_lookup = {str(rid): row for rid, (_, row) in zip(ids, df.iterrows())}
        n_lookup = dict(zip(ids, ns))
        for rid, txt in zip(active_ids, texts):
            if txt:
                a = apply_parsers(rid, row_lookup[rid], n_lookup[rid], txt)
                samples[rid].append(a)
        for rid in ids:
            if len(samples[rid]) >= 3:
                greedy = samples[rid][0]
                n = n_lookup[rid] if rid in n_lookup else len(greedy)
                preds[rid] = enforce_bijection([
                    vote_answer([s[k] for s in samples[rid] if k < len(s)],
                                anchor=greedy[k] if k < len(greedy) else None)
                    for k in range(n)
                ])
        write_submission(OUT_CSV, ids, preds, explanations)
        log(f"  voted over up to {n_extra + 1} samples/row -> written")
        report_dev_score(preds)

    # 5. Explanations for the jury track.
    if WANT_EXPLANATION and time_left() > 60:
        log(f"generating explanations ({time_left():.0f}s left)")
        ex_prompts = []
        for (_, r), i in zip(df.iterrows(), ids):
            msgs = [{"role": "system", "content": EXPLAIN_SYSTEM},
                    {"role": "user", "content": build_explain_prompt(r, preds[i])}]
            ex_prompts.append(tok.apply_chat_template(
                msgs, tokenize=False, add_generation_prompt=True))
        ex_texts = generate(ex_prompts, max_new=200, sample=False)
        for i, e in zip(ids, ex_texts):
            e = re.sub(r"\s+", " ", (e or "").strip())
            if e:
                explanations[i] = e[:1200]
        write_submission(OUT_CSV, ids, preds, explanations)
        log("explanations written")

    # 6. Final integrity pass: guarantee shape and non-blank answers no
    # matter what happened above.
    n_lookup = dict(zip(ids, ns))
    malformed = [i for i in ids if len(preds[i]) != n_lookup[i]
                 or any(not str(x).strip() for x in preds[i])]
    if malformed:
        log(f"repairing {len(malformed)} malformed rows")
        for i in malformed:
            preds[i] = squeeze_to_n(
                [x for x in preds[i] if str(x).strip()], n_lookup[i], fallbacks[i])
        write_submission(OUT_CSV, ids, preds, explanations)

    log(f"DONE. {len(ids)} rows, {sum(len(v) for v in preds.values())} answers, "
        f"{time.time() - START:.0f}s elapsed")


if __name__ == "__main__":
    main()