| """ |
| 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() |
|
|