Spatial9/GravityLLM
0
1import argparse2import json3import re4from pathlib import Path5from typing import Dict, Tuple6 7import torch8from datasets import load_dataset9from jsonschema import Draft7Validator10from peft import AutoPeftModelForCausalLM11from transformers import AutoModelForCausalLM, AutoTokenizer12 13SYSTEM_PREFIX = (14 "You are GravityLLM, a Spatial9 scene generation model. "15 "Given music constraints and stem features, output ONLY valid Spatial9Scene JSON. "16 "Do not return markdown. Do not explain your answer.\n\n"17)18 19 20def parse_args() -> argparse.Namespace:21 parser = argparse.ArgumentParser(description="Evaluate GravityLLM outputs on a JSONL validation set.")22 parser.add_argument("--model_dir", type=str, required=True)23 parser.add_argument("--data_file", type=str, default="data/valid.jsonl")24 parser.add_argument("--schema_path", type=Path, default=Path("schemas/scene.schema.json"))25 parser.add_argument("--max_new_tokens", type=int, default=900)26 parser.add_argument("--temperature", type=float, default=0.2)27 parser.add_argument("--top_p", type=float, default=0.9)28 parser.add_argument("--limit", type=int, default=0, help="0 means evaluate all rows.")29 parser.add_argument("--report_path", type=Path, default=Path("reports/eval_report.json"))30 return parser.parse_args()31 32 33def load_model_and_tokenizer(model_dir: str):34 tokenizer = AutoTokenizer.from_pretrained(model_dir, use_fast=True, trust_remote_code=True)35 if tokenizer.pad_token is None:36 tokenizer.pad_token = tokenizer.eos_token37 38 try:39 model = AutoPeftModelForCausalLM.from_pretrained(40 model_dir,41 torch_dtype=torch.bfloat16 if torch.cuda.is_available() else None,42 device_map="auto" if torch.cuda.is_available() else None,43 trust_remote_code=True,44 )45 except Exception:46 model = AutoModelForCausalLM.from_pretrained(47 model_dir,48 torch_dtype=torch.bfloat16 if torch.cuda.is_available() else None,49 device_map="auto" if torch.cuda.is_available() else None,50 trust_remote_code=True,51 )52 model.eval()53 return model, tokenizer54 55 56def format_prompt(raw_prompt: str) -> str:57 raw_prompt = raw_prompt.strip()58 if raw_prompt.lower().startswith("gravityllm:"):59 raw_prompt = raw_prompt.split(":", 1)[1].strip()60 return SYSTEM_PREFIX + raw_prompt + "\n\nOUTPUT:\n"61 62 63def extract_first_json(text: str) -> str:64 match = re.search(r"\{.*\}", text, flags=re.DOTALL)65 return match.group(0).strip() if match else text.strip()66 67 68def validate_schema(schema, output_text: str) -> Tuple[bool, Dict]:69 data = json.loads(output_text)70 validator = Draft7Validator(schema)71 errors = sorted(validator.iter_errors(data), key=lambda e: list(e.path))72 return len(errors) == 0, data73 74 75def check_budget(input_payload: Dict, scene_payload: Dict) -> bool:76 max_objects = input_payload.get("max_objects")77 if max_objects is None:78 return True79 return len(scene_payload.get("objects", [])) <= max_objects80 81 82def check_anchor_rules(input_payload: Dict, scene_payload: Dict) -> bool:83 objects = {obj["class"]: obj for obj in scene_payload.get("objects", [])}84 for rule in input_payload.get("rules", []):85 if rule.get("type") != "anchor":86 continue87 klass = rule.get("track_class")88 obj = objects.get(klass)89 if obj is None:90 return False91 for field in ["az_deg", "el_deg", "dist_m"]:92 if float(obj[field]) != float(rule[field]):93 return False94 return True95 96 97def generate_scene(model, tokenizer, prompt_text: str, max_new_tokens: int, temperature: float, top_p: float) -> str:98 inputs = tokenizer(prompt_text, return_tensors="pt")99 if torch.cuda.is_available():100 inputs = {k: v.to(model.device) for k, v in inputs.items()}101 102 with torch.no_grad():103 outputs = model.generate(104 **inputs,105 max_new_tokens=max_new_tokens,106 do_sample=True,107 temperature=temperature,108 top_p=top_p,109 eos_token_id=tokenizer.eos_token_id,110 pad_token_id=tokenizer.pad_token_id,111 )112 113 decoded = tokenizer.decode(outputs[0], skip_special_tokens=True)114 prompt_prefix = tokenizer.decode(inputs["input_ids"][0], skip_special_tokens=True)115 raw_completion = decoded[len(prompt_prefix):].strip()116 return extract_first_json(raw_completion)117 118 119def main() -> None:120 args = parse_args()121 schema = json.loads(args.schema_path.read_text(encoding="utf-8"))122 ds = load_dataset("json", data_files=args.data_file, split="train")123 if args.limit > 0:124 ds = ds.select(range(min(args.limit, len(ds))))125 126 model, tokenizer = load_model_and_tokenizer(args.model_dir)127 128 total = len(ds)129 parse_ok = 0130 schema_ok = 0131 budget_ok = 0132 anchor_ok = 0133 samples = []134 135 for row in ds:136 prompt_text = format_prompt(row["prompt"])137 generated = generate_scene(model, tokenizer, prompt_text, args.max_new_tokens, args.temperature, args.top_p)138 139 sample_report = {"prompt": row["prompt"], "generated": generated}140 try:141 gen_data = json.loads(generated)142 parse_ok += 1143 valid, gen_scene = validate_schema(schema, generated)144 if valid:145 schema_ok += 1146 # Reconstruct input payload from prompt for simple rule checks.147 prompt_payload_text = row["prompt"].split("INPUT:\n", 1)[1]148 input_payload = json.loads(prompt_payload_text)149 if check_budget(input_payload, gen_scene):150 budget_ok += 1151 if check_anchor_rules(input_payload, gen_scene):152 anchor_ok += 1153 sample_report["schema_valid"] = True154 sample_report["budget_pass"] = check_budget(input_payload, gen_scene)155 sample_report["anchor_pass"] = check_anchor_rules(input_payload, gen_scene)156 else:157 sample_report["schema_valid"] = False158 except Exception as exc:159 sample_report["error"] = str(exc)160 161 samples.append(sample_report)162 163 report = {164 "examples": total,165 "json_parse_rate": round(parse_ok / total, 4) if total else 0.0,166 "schema_valid_rate": round(schema_ok / total, 4) if total else 0.0,167 "budget_pass_rate": round(budget_ok / total, 4) if total else 0.0,168 "anchor_pass_rate": round(anchor_ok / total, 4) if total else 0.0,169 "samples": samples[:10],170 }171 172 args.report_path.parent.mkdir(parents=True, exist_ok=True)173 args.report_path.write_text(json.dumps(report, indent=2), encoding="utf-8")174 print(json.dumps(report, indent=2))175 176 177if __name__ == "__main__":178 main()179 