File size: 3,540 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
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
"""

Inference and generation script for Retriever500M.



Loads a checkpoint and generates text to verify the model has learned

code structure during base pretraining.



Usage:

  python src/generate.py [--checkpoint PATH] [--prompt "text"] [--tokens N]

"""

import argparse
import os
import sys

import torch

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

from model import ModelConfig, Retriever500M
from tokenizers import Tokenizer

PROJECT_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
CHECKPOINT_DIR = os.path.join(PROJECT_DIR, "checkpoints")
TOKENIZER_PATH = os.path.join(PROJECT_DIR, "tokenizer", "tokenizer.json")


def load_model(checkpoint_path: str, device: torch.device) -> tuple[Retriever500M, ModelConfig]:
    """Load model from checkpoint."""
    print(f"Loading checkpoint: {checkpoint_path}")
    ckpt = torch.load(checkpoint_path, map_location=device, weights_only=False)

    config = ModelConfig(**ckpt["config"])
    model = Retriever500M(config).to(device)
    model.load_state_dict(ckpt["model_state_dict"])
    model.eval()

    print(f"  Step: {ckpt.get('step', '?')}")
    print(f"  Loss: {ckpt.get('loss', '?')}")
    print(f"  Params: {model.count_parameters() / 1e6:.1f}M")

    return model, config


def generate(

    model: Retriever500M,

    tokenizer: Tokenizer,

    prompt: str,

    max_new_tokens: int = 128,

    temperature: float = 0.8,

    top_k: int = 50,

    device: torch.device = None,

) -> str:
    """Generate text from a prompt."""
    if device is None:
        device = next(model.parameters()).device

    # Encode prompt
    encoded = tokenizer.encode(prompt)
    input_ids = torch.tensor([encoded.ids], dtype=torch.long, device=device)

    # Generate
    with torch.no_grad():
        output_ids = model.generate(
            input_ids,
            max_new_tokens=max_new_tokens,
            temperature=temperature,
            top_k=top_k,
        )

    # Decode
    output_text = tokenizer.decode(output_ids[0].tolist())
    return output_text


def main():
    parser = argparse.ArgumentParser(description="Generate text with Retriever500M")
    parser.add_argument("--checkpoint", type=str, default=os.path.join(CHECKPOINT_DIR, "latest.pt"))
    parser.add_argument("--prompt", type=str, default="def fibonacci(n):\n    ", help="Generation prompt")
    parser.add_argument("--tokens", type=int, default=128, help="Max new tokens")
    parser.add_argument("--temperature", type=float, default=0.8)
    parser.add_argument("--top_k", type=int, default=50)
    args = parser.parse_args()

    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    print(f"Device: {device}")

    # Load tokenizer
    tokenizer = Tokenizer.from_file(TOKENIZER_PATH)

    # Load model
    model, config = load_model(args.checkpoint, device)

    # Generate
    prompts = [
        args.prompt,
        "def quicksort(arr):\n    ",
        "function fetchData(url) {\n  ",
        "import torch\nimport torch.nn as nn\n\nclass Model(nn.Module):\n    ",
        "fn main() {\n    println!",
    ]

    print("\n" + "=" * 60)
    print("GENERATION SAMPLES")
    print("=" * 60)

    for prompt in prompts:
        print(f"\n--- Prompt: {prompt!r} ---")
        text = generate(model, tokenizer, prompt, args.tokens, args.temperature, args.top_k, device)
        print(text)
        print("-" * 60)


if __name__ == "__main__":
    main()