File size: 4,470 Bytes
d3ab6f5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Minimal reference sampler for the public Vortex Alpha release."""

from __future__ import annotations

import argparse
import json
from pathlib import Path

import torch
from safetensors.torch import load_file

from vortex_model import VortexConfig, VortexForCausalLM


def load_config(path: Path) -> VortexConfig:
    data = json.loads(path.read_text())
    return VortexConfig(
        vocab_size=int(data["vocab_size"]),
        max_seq_len=int(data["max_position_embeddings"]),
        n_layer=int(data["num_hidden_layers"]),
        n_embd=int(data["hidden_size"]),
        n_head=int(data["num_attention_heads"]),
        n_kv_head=int(data["num_key_value_heads"]),
        head_dim=int(data["head_dim"]),
        intermediate_size=int(data["intermediate_size"]),
        rope_theta=float(data["rope_theta"]),
        norm_eps=float(data["rms_norm_eps"]),
        logits_chunk_tokens=8192,
        gradient_checkpointing=False,
        use_transformer_engine=False,
        attn_input_format="bshd",
    )


def choose_device(requested: str) -> torch.device:
    if requested != "auto":
        return torch.device(requested)
    if torch.cuda.is_available():
        return torch.device("cuda")
    if getattr(torch.backends, "mps", None) is not None and torch.backends.mps.is_available():
        return torch.device("mps")
    return torch.device("cpu")


def sample_next(logits: torch.Tensor, temperature: float, top_p: float) -> torch.Tensor:
    if temperature <= 0:
        return logits.argmax(dim=-1, keepdim=True)
    logits = logits / temperature
    if 0 < top_p < 1:
        sorted_logits, sorted_indices = torch.sort(logits, descending=True, dim=-1)
        sorted_probs = torch.softmax(sorted_logits, dim=-1)
        cumulative = sorted_probs.cumsum(dim=-1)
        remove = cumulative - sorted_probs > top_p
        sorted_logits = sorted_logits.masked_fill(remove, float("-inf"))
        logits = torch.full_like(logits, float("-inf"))
        logits.scatter_(dim=-1, index=sorted_indices, src=sorted_logits)
    return torch.multinomial(torch.softmax(logits, dim=-1), num_samples=1)


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--weights", type=Path, default=Path("model.safetensors"))
    parser.add_argument("--config", type=Path, default=Path("config.json"))
    parser.add_argument("--tokenizer", type=Path, default=Path("tokenizer.model"))
    parser.add_argument("--prompt", required=True)
    parser.add_argument("--chat", action="store_true")
    parser.add_argument(
        "--system",
        default="You are a helpful assistant. Answer clearly and say when information is missing.",
    )
    parser.add_argument("--max-new-tokens", type=int, default=128)
    parser.add_argument("--temperature", type=float, default=0.0)
    parser.add_argument("--top-p", type=float, default=0.95)
    parser.add_argument("--device", default="auto")
    args = parser.parse_args()

    import sentencepiece as spm

    device = choose_device(args.device)
    dtype = torch.bfloat16 if device.type == "cuda" else torch.float32
    config = load_config(args.config)
    model = VortexForCausalLM(config).to(device=device, dtype=dtype)
    state = load_file(str(args.weights), device="cpu")
    model.load_state_dict(state, strict=True)
    model.eval()

    prompt = args.prompt
    if args.chat:
        prompt = f"[SYSTEM]\n{args.system}\n</s>\n[USER]\n{prompt}\n</s>\n[ASSISTANT]\n"
    tokenizer = spm.SentencePieceProcessor(model_file=str(args.tokenizer))
    ids = tokenizer.encode(prompt, out_type=int)
    if not ids:
        raise ValueError("prompt encoded to zero tokens")
    if len(ids) >= config.max_seq_len:
        raise ValueError("prompt reaches the configured context limit")

    input_ids = torch.tensor([ids], dtype=torch.long, device=device)
    generated: list[int] = []
    with torch.inference_mode():
        for _ in range(max(0, args.max_new_tokens)):
            logits, _ = model(input_ids)
            next_id = sample_next(logits[:, -1, :].float(), args.temperature, args.top_p)
            token_id = int(next_id.item())
            if token_id == tokenizer.eos_id():
                break
            generated.append(token_id)
            input_ids = torch.cat((input_ids, next_id), dim=1)
            if input_ids.shape[1] >= config.max_seq_len:
                break
    print(tokenizer.decode(generated, out_type=str), end="")


if __name__ == "__main__":
    main()