Download src/verify_pipeline.py from Reizxn/makeitwork: direct link, hf CLI and curl.
- Browser
- Download file 3.33 kB
-
https://huggingface.co/Reizxn/makeitwork/resolve/main/src/verify_pipeline.py
- Command line
-
hf download hf://Reizxn/makeitwork/src/verify_pipeline.py
-
curl -L -o verify_pipeline.py https://huggingface.co/Reizxn/makeitwork/resolve/main/src/verify_pipeline.py
3.33 kB
| """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() | |