Spatial9/GravityLLM
0
1import argparse2import json3import re4from pathlib import Path5 6import torch7from jsonschema import Draft7Validator8from peft import AutoPeftModelForCausalLM9from transformers import AutoModelForCausalLM, AutoTokenizer10 11SYSTEM_PREFIX = (12 "You are GravityLLM, a Spatial9 scene generation model. "13 "Given music constraints and stem features, output ONLY valid Spatial9Scene JSON. "14 "Do not return markdown. Do not explain your answer.\n\n"15)16 17 18def parse_args() -> argparse.Namespace:19 parser = argparse.ArgumentParser(description="Run GravityLLM inference on a Spatial9 constraint payload.")20 parser.add_argument("--model_dir", type=str, required=True, help="Path or Hub repo id for trained model or adapter.")21 parser.add_argument("--input_json", type=Path, required=True)22 parser.add_argument("--schema_path", type=Path, default=Path("schemas/scene.schema.json"))23 parser.add_argument("--output_json", type=Path, default=None)24 parser.add_argument("--max_new_tokens", type=int, default=900)25 parser.add_argument("--temperature", type=float, default=0.35)26 parser.add_argument("--top_p", type=float, default=0.9)27 parser.add_argument("--validate", action="store_true")28 return parser.parse_args()29 30 31def load_model_and_tokenizer(model_dir: str):32 tokenizer = AutoTokenizer.from_pretrained(model_dir, use_fast=True, trust_remote_code=True)33 if tokenizer.pad_token is None:34 tokenizer.pad_token = tokenizer.eos_token35 36 model = None37 try:38 model = AutoPeftModelForCausalLM.from_pretrained(39 model_dir,40 torch_dtype=torch.bfloat16 if torch.cuda.is_available() else None,41 device_map="auto" if torch.cuda.is_available() else None,42 trust_remote_code=True,43 )44 except Exception:45 model = AutoModelForCausalLM.from_pretrained(46 model_dir,47 torch_dtype=torch.bfloat16 if torch.cuda.is_available() else None,48 device_map="auto" if torch.cuda.is_available() else None,49 trust_remote_code=True,50 )51 model.eval()52 return model, tokenizer53 54 55def extract_first_json(text: str) -> str:56 match = re.search(r"\{.*\}", text, flags=re.DOTALL)57 return match.group(0).strip() if match else text.strip()58 59 60def validate_output(schema_path: Path, output_text: str):61 schema = json.loads(schema_path.read_text(encoding="utf-8"))62 data = json.loads(output_text)63 validator = Draft7Validator(schema)64 errors = sorted(validator.iter_errors(data), key=lambda e: list(e.path))65 return data, errors66 67 68def main() -> None:69 args = parse_args()70 payload = json.loads(args.input_json.read_text(encoding="utf-8"))71 72 model, tokenizer = load_model_and_tokenizer(args.model_dir)73 prompt = SYSTEM_PREFIX + "INPUT:\n" + json.dumps(payload, ensure_ascii=False, indent=2) + "\n\nOUTPUT:\n"74 75 inputs = tokenizer(prompt, return_tensors="pt")76 if torch.cuda.is_available():77 inputs = {k: v.to(model.device) for k, v in inputs.items()}78 79 with torch.no_grad():80 outputs = model.generate(81 **inputs,82 max_new_tokens=args.max_new_tokens,83 do_sample=True,84 temperature=args.temperature,85 top_p=args.top_p,86 eos_token_id=tokenizer.eos_token_id,87 pad_token_id=tokenizer.pad_token_id,88 )89 90 decoded = tokenizer.decode(outputs[0], skip_special_tokens=True)91 prompt_prefix = tokenizer.decode(inputs["input_ids"][0], skip_special_tokens=True)92 raw_completion = decoded[len(prompt_prefix):].strip()93 json_text = extract_first_json(raw_completion)94 95 if args.validate:96 try:97 _, errors = validate_output(args.schema_path, json_text)98 if errors:99 print("Validation: INVALID")100 for err in errors[:20]:101 path = ".".join(str(p) for p in err.path)102 print(f"- {path}: {err.message}")103 else:104 print("Validation: VALID")105 except Exception as exc:106 print(f"Validation failed: {exc}")107 108 if args.output_json:109 args.output_json.parent.mkdir(parents=True, exist_ok=True)110 args.output_json.write_text(json_text + "\n", encoding="utf-8")111 print(f"Wrote output to {args.output_json}")112 113 print(json_text)114 115 116if __name__ == "__main__":117 main()118 