CoolFace
Modelpublic

Spatial9/GravityLLM

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