CoolFace
Modelpublic

Spatial9/GravityLLM

sourceHugging Faceapache-2.0updated 7mo agoView on Hugging Face
0likes
infer.py118 linesDownload Raw Back to root
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