File size: 3,331 Bytes
3738348
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Verify the full search agent pipeline end-to-end."""
import json
import os
import sys

sys.path.insert(0, os.path.join(os.path.dirname(os.path.abspath(__file__))))

def main():
    # 1. Verify chunks load
    chunks = [json.loads(l) for l in open("data/chunks.jsonl", encoding="utf-8")]
    print(f"1. Chunks: {len(chunks):,} loaded OK")

    # 2. Verify agent tokenizer
    from tokenizers import Tokenizer
    tok = Tokenizer.from_file("tokenizer/tokenizer_agent.json")
    print(f"2. Agent tokenizer: vocab={tok.get_vocab_size()} OK")

    # 3. Verify special tokens
    st = json.load(open("tokenizer/special_tokens.json"))
    ids = list(st["token_ids"].values())
    print(f"3. Special tokens: {len(st['special_tokens'])} tokens, IDs {ids} OK")

    # 4. Verify SFT traces load and format correctly
    from sft import format_trace_to_tokens, load_special_tokens
    load_special_tokens()
    traces = [json.loads(l) for l in open("data/sft_traces.jsonl", encoding="utf-8")]
    print(f"4. SFT traces: {len(traces):,} loaded")

    # Format one trace
    ids_arr, mask_arr = format_trace_to_tokens(traces[0], tok, max_seq_len=768)
    loss_tokens = int(mask_arr.sum())
    total_tokens = len(ids_arr)
    pct = loss_tokens / total_tokens * 100
    print(f"   Sample trace: {total_tokens} tokens, {loss_tokens} with loss ({pct:.0f}%)")

    # 5. Verify gold traces
    gold = [json.loads(l) for l in open("data/gold_traces.jsonl", encoding="utf-8")]
    print(f"5. Gold traces: {len(gold)} loaded OK")

    # 6. Verify model accepts new vocab size
    from model import ModelConfig, Retriever500M
    import torch
    config = ModelConfig(vocab_size=32009, d_model=1280, n_layers=23, n_heads=20, d_ff=3456, max_seq_len=768)
    model = Retriever500M(config)
    print(f"6. Model with vocab=32009: {model.count_parameters()/1e6:.1f}M params OK")

    # 7. Verify model can load old checkpoint with embedding resize
    ckpt = torch.load("checkpoints/latest.pt", map_location="cpu", weights_only=False)
    old_vocab = ckpt["config"]["vocab_size"]
    needs_resize = old_vocab != 32009
    print(f"7. Old checkpoint vocab={old_vocab}, new vocab=32009, resize needed={needs_resize}")

    # 8. Verify retriever
    from search_agent import KeywordRetriever
    retr = KeywordRetriever("data/chunks.jsonl")
    results = retr.search("ngx_reusable_connection", top_k=3)
    print(f"8. Retriever: search returned {len(results)} results OK")
    if results:
        print(f"   Top result: {results[0]['name']} (score={results[0]['score']:.2f})")

    # 9. Verify SFT data formatting produces valid tokens
    from sft import get_sft_batch
    dataset = []
    for trace in traces[:100]:
        ids, mask = format_trace_to_tokens(trace, tok, max_seq_len=768)
        if len(ids) > 10:
            dataset.append((ids, mask))
    input_ids, targets, loss_mask = get_sft_batch(dataset, batch_size=2, seq_len=768, device=torch.device("cpu"))
    print(f"9. SFT batch: input_ids={input_ids.shape}, targets={targets.shape}, loss_mask={loss_mask.shape}")
    print(f"   Loss mask coverage: {loss_mask.sum().item()}/{loss_mask.numel()} ({loss_mask.float().mean()*100:.0f}%)")

    print()
    print("ALL CHECKS PASSED")


if __name__ == "__main__":
    main()