CoolFace
Modelpublic

Chaman1234/Sparse-AST-BWM

sourceHugging Faceapache-2.0updated 8d agoView on Hugging Face
0likes1.7kdownloads
generate.py67 linesDownload Raw Back to root
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