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