Chaman1234/Sparse-AST-BWM
01.7k
1"""
2Unified Autoregressive Generation & MoE Routing CLI for Sparse-AST / BWM
3"""
4
5import os
6import argparse
7import torch
8import torch.nn.functional as F
9from model import SparseAST, TopKSparseASTEnsemble
10
11def generate_backbone(model, prompt, max_new_tokens=40, temperature=0.6, top_k=40):
12 model.eval()
13 device = next(model.parameters()).device
14 tokens = list(prompt.encode("utf-8", "ignore"))
15
16 with torch.no_grad():
17 for _ in range(max_new_tokens):
18 x = torch.tensor([tokens[-model.seq_len:]], dtype=torch.long, device=device)
19 logits = model(x)[:, -1, :]
20
21 if temperature <= 0.05:
22 next_tok = int(logits.argmax(dim=-1).item())
23 else:
24 scaled = logits / temperature
25 if top_k > 0:
26 v, _ = torch.topk(scaled, min(top_k, scaled.size(-1)))
27 scaled[scaled < v[:, [-1]]] = -float("inf")
28 probs = F.softmax(scaled, dim=-1)
29 next_tok = int(torch.multinomial(probs, num_samples=1).item())
30 tokens.append(next_tok)
31
32 return bytes([t for t in tokens if t < 256]).decode("utf-8", errors="replace")
33
34def main():
35 parser = argparse.ArgumentParser(description="Generate Blender Python code using Sparse-AST / BWM")
36 parser.add_argument("--model", type=str, default="100M-32", help="Model variant: 3M-32, 10M-32, 100M-32, 200M-32, moe")
37 parser.add_argument("--prompt", type=str, default="import bpy\n# Create primitive cube\nbpy.ops.mesh.", help="Prompt text")
38 parser.add_argument("--tokens", type=int, default=40, help="Max tokens to generate")
39 parser.add_argument("--temperature", type=float, default=0.6, help="Sampling temperature")
40 args = parser.parse_args()
41
42 repo_dir = os.path.dirname(os.path.abspath(__file__))
43 device = "cuda" if torch.cuda.is_available() else "cpu"
44
45 print("=" * 80)
46 print(f" SPARSE-AST / BWM CODE GENERATION: Model [{args.model}] on [{device}]")
47 print("=" * 80)
48
49 if args.model.lower() in ["moe", "topk-moe"]:
50 ensemble = TopKSparseASTEnsemble.from_pretrained(repo_dir, subfolder="TopK-MoE", device=device)
51 x = torch.tensor([list(args.prompt.encode("utf-8", "ignore"))], dtype=torch.long, device=device)
52 _, topk_indices, gate_weights = ensemble(x, k=2)
53 print(f"\nPrompt: {repr(args.prompt)}")
54 print("Top-2 Experts selected for next token:")
55 for idx, w in zip(topk_indices[0, -1].tolist(), gate_weights[0, -1].tolist()):
56 exp_name = ensemble.expert_names[idx] if idx < len(ensemble.expert_names) else f"Expert {idx}"
57 print(f" {exp_name}: Weight {w:.4f}")
58 else:
59 sub = args.model if args.model.endswith("-32") or args.model.endswith("-512") else f"{args.model}-32"
60 model = SparseAST.from_pretrained(repo_dir, subfolder=sub, device=device)
61 print(f"\nPrompt:\n{args.prompt}")
62 completion = generate_backbone(model, args.prompt, max_new_tokens=args.tokens, temperature=args.temperature)
63 print(f"\nGenerated Completion:\n{completion}")
64
65if __name__ == "__main__":
66 main()
67 