File size: 3,671 Bytes
e38f140
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
TrinityX inference — load the model and generate responses.
Requires: HF_TOKEN environment variable (for LLaMA-2 access).

Usage:
    python inference.py --prompt "What is climate change?"
    python inference.py --interactive
"""
import os, json, torch, argparse, yaml
from models.base_model import load_base_model, load_tokenizer
from models.mocae_model import TrinityXModel


def load_trinityX(model_dir: str = ".", hf_token: str = None):
    token = hf_token or os.environ.get("HF_TOKEN")

    with open(os.path.join(model_dir, "config.yaml")) as f:
        cfg = yaml.safe_load(f)

    mocae_config = {
        "num_experts":      cfg.get("num_experts", 3),
        "router_hidden_dim": cfg.get("router_hidden_dim", 128),
        "router_output_dim": cfg.get("router_output_dim", 64),
        "temperature":      cfg.get("temperature", 0.7),
        "dropout_rate":     cfg.get("dropout_rate", 0.1),
    }

    print("Loading base model...")
    tokenizer  = load_tokenizer(cfg["backbone"], hf_token=token)
    tokenizer.padding_side = "left"
    base_model = load_base_model(cfg["backbone"], precision=cfg.get("precision", "bfloat16"),
                                 device_map="auto", hf_token=token, freeze=True)

    adapter_paths = [
        os.path.join(model_dir, "expert_helpfulness"),
        os.path.join(model_dir, "expert_harmlessness"),
        os.path.join(model_dir, "expert_honesty"),
    ]

    print("Loading TrinityX MoCaE...")
    model = TrinityXModel.load_pretrained(
        save_dir=os.path.join(model_dir, "trinityX_final"),
        base_model=base_model,
        adapter_paths=adapter_paths,
        mocae_config=mocae_config,
        lora_rank=cfg.get("lora_rank", 16),
        lora_alpha=cfg.get("lora_alpha", 32),
    )
    model.eval()
    dev = next(base_model.parameters()).device
    print(f"TrinityX ready on {dev}")
    return model, tokenizer, dev


def generate(model, tokenizer, dev, prompt: str, max_new_tokens: int = 200):
    formatted = f"[INST] {prompt} [/INST]"
    inputs = tokenizer(formatted, return_tensors="pt", truncation=True,
                       max_length=512).to(dev)
    with torch.no_grad():
        out = model.base_model.generate(
            **inputs,
            max_new_tokens=max_new_tokens,
            do_sample=False,
            repetition_penalty=1.1,
            pad_token_id=tokenizer.pad_token_id,
            eos_token_id=tokenizer.eos_token_id,
        )
    pl = inputs["input_ids"].shape[1]
    return tokenizer.decode(out[0][pl:], skip_special_tokens=True).strip()


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--prompt", type=str, default=None)
    parser.add_argument("--interactive", action="store_true")
    parser.add_argument("--model_dir", type=str, default=".")
    parser.add_argument("--max_tokens", type=int, default=200)
    args = parser.parse_args()

    model, tokenizer, dev = load_trinityX(args.model_dir)

    if args.interactive:
        print("\nTrinityX — type your question (Ctrl+C to exit)\n")
        while True:
            try:
                prompt = input("You: ").strip()
                if not prompt:
                    continue
                response = generate(model, tokenizer, dev, prompt, args.max_tokens)
                print(f"TrinityX: {response}\n")
            except KeyboardInterrupt:
                print("\nGoodbye!")
                break
    elif args.prompt:
        response = generate(model, tokenizer, dev, args.prompt, args.max_tokens)
        print(f"TrinityX: {response}")
    else:
        print("Use --prompt 'your question' or --interactive")


if __name__ == "__main__":
    main()