| """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():
|
|
|
| chunks = [json.loads(l) for l in open("data/chunks.jsonl", encoding="utf-8")]
|
| print(f"1. Chunks: {len(chunks):,} loaded OK")
|
|
|
|
|
| from tokenizers import Tokenizer
|
| tok = Tokenizer.from_file("tokenizer/tokenizer_agent.json")
|
| print(f"2. Agent tokenizer: vocab={tok.get_vocab_size()} OK")
|
|
|
|
|
| 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")
|
|
|
|
|
| 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")
|
|
|
|
|
| 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}%)")
|
|
|
|
|
| gold = [json.loads(l) for l in open("data/gold_traces.jsonl", encoding="utf-8")]
|
| print(f"5. Gold traces: {len(gold)} loaded OK")
|
|
|
|
|
| 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")
|
|
|
|
|
| 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}")
|
|
|
|
|
| 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})")
|
|
|
|
|
| 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()
|
|
|