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