AI-Puppet-Theater / finetune /scripts /validate_actor_sft.py
ShubhamSetia's picture
finetune: add Actor SFT v1 audit and eval hardening
34e3cf7
Raw
History Blame Contribute Delete
13.6 kB
#!/usr/bin/env python3
"""Validate AI Puppet Theater Actor SFT JSONL files."""
from __future__ import annotations
import argparse
import json
from collections import Counter
from pathlib import Path
from typing import Any
REQUIRED_TOP_LEVEL = {"id", "source_mix", "row_type", "messages"}
OPTIONAL_TOP_LEVEL = {"source_dataset", "transformation"}
ALLOWED_TOP_LEVEL = REQUIRED_TOP_LEVEL | OPTIONAL_TOP_LEVEL
ALLOWED_SOURCE_DATASETS = {
"G-reen/TheatreLM-v2.1-Characters",
"practical-dreamer/RPGPT_PublicDomain-alpaca",
}
ALLOWED_ROW_TYPES = {
"normal_reaction",
"prop_inspection",
"oracle_consult",
"lighting_change",
"memory_callback",
"secret_hint_or_reveal",
"finale",
"comedic_confusion",
}
OUTPUT_FIELDS = [
"intent",
"line",
"emotion",
"gesture",
"stage_effect",
"memory_update",
"tool_request",
]
ALLOWED_INTENTS = {
"react_to_event",
"clarify_problem",
"inspect_prop",
"consult_oracle",
"change_lighting",
"recall_memory",
"hint_secret",
"reveal_secret",
"deliver_finale",
"comic_confusion",
}
ALLOWED_TOOLS = {"inspect_prop", "consult_stage_oracle", "change_lighting"}
TOOL_ARGS = {
"inspect_prop": {"prop"},
"consult_stage_oracle": {"question"},
"change_lighting": {"mood"},
}
V1_TOOL_ARGS = {
"inspect_prop": {"prop"},
"consult_stage_oracle": {"question"},
"change_lighting": {"mood"},
}
V1_BANNED_ASSISTANT_FIELDS = {
"memory_record",
"memory_effect",
"recent_transcript",
"show_state",
"held_props",
"mood",
"name",
"latest_prop",
"latest_audience_action",
"tool_results",
}
BANNED_TERMS = {
"fuck",
"shit",
"bitch",
"asshole",
"bloodbath",
"gore",
"dismember",
"suicide",
}
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("path", type=Path)
parser.add_argument("--max-errors", type=int, default=25)
args = parser.parse_args()
stats = validate_file(args.path, args.max_errors)
print(f"file: {args.path}")
print(f"total rows: {stats['total']}")
print(f"valid rows: {stats['valid']}")
print(f"invalid rows: {stats['invalid']}")
print_distribution("row_type distribution", stats["row_types"])
print_distribution("tool_request distribution", stats["tools"])
stem = args.path.stem
if stem.endswith("_train"):
split_stem = stem[:-6]
elif stem.endswith("_val"):
split_stem = stem[:-4]
else:
split_stem = stem
train_path = args.path.with_name(f"{split_stem}_train.jsonl")
val_path = args.path.with_name(f"{split_stem}_val.jsonl")
if train_path.exists() or val_path.exists():
print(f"train rows: {count_lines(train_path) if train_path.exists() else 0}")
print(f"val rows: {count_lines(val_path) if val_path.exists() else 0}")
if stats["errors"]:
print("errors:")
for error in stats["errors"]:
print(f"- {error}")
if stats["invalid"]:
raise SystemExit(1)
def validate_file(path: Path, max_errors: int) -> dict[str, Any]:
stats: dict[str, Any] = {
"total": 0,
"valid": 0,
"invalid": 0,
"row_types": Counter(),
"tools": Counter(),
"errors": [],
}
with path.open("r", encoding="utf-8") as handle:
for line_number, line in enumerate(handle, start=1):
if not line.strip():
continue
stats["total"] += 1
try:
row = json.loads(line)
except json.JSONDecodeError as exc:
add_error(stats, max_errors, line_number, f"row JSON parse failed: {exc}")
continue
errors = validate_row(row, dataset_version_for(path, row))
if errors:
stats["invalid"] += 1
for error in errors:
add_error(stats, max_errors, line_number, error)
continue
assistant = json.loads(row["messages"][2]["content"])
stats["valid"] += 1
stats["row_types"][row["row_type"]] += 1
tool_request = assistant["tool_request"]
stats["tools"][tool_request["tool"] if tool_request else "none"] += 1
return stats
def dataset_version_for(path: Path, row: Any) -> str:
if "v1" in path.name:
return "v1"
if isinstance(row, dict):
row_id = str(row.get("id", ""))
source_mix = row.get("source_mix", [])
if row_id.startswith("actor-sft-v1-") or "synthetic_v1" in source_mix:
return "v1"
return "v0"
def validate_row(row: Any, version: str = "v0") -> list[str]:
errors: list[str] = []
if not isinstance(row, dict):
return ["row must be an object"]
unknown_top_level = set(row) - ALLOWED_TOP_LEVEL
missing_top_level = REQUIRED_TOP_LEVEL - set(row)
if unknown_top_level:
errors.append(f"unknown top-level keys: {sorted(unknown_top_level)}")
if missing_top_level:
errors.append(f"missing top-level keys: {sorted(missing_top_level)}")
return errors
if not isinstance(row["id"], str) or not row["id"].strip():
errors.append("id must be a non-empty string")
if row["row_type"] not in ALLOWED_ROW_TYPES:
errors.append(f"unknown row_type: {row['row_type']!r}")
if not isinstance(row["source_mix"], list) or not all(isinstance(item, str) and item for item in row["source_mix"]):
errors.append("source_mix must be a non-empty list of strings")
if "source_dataset" in row:
if row["source_dataset"] not in ALLOWED_SOURCE_DATASETS:
errors.append(f"source_dataset must be one of {sorted(ALLOWED_SOURCE_DATASETS)}")
if row["source_dataset"] == "G-reen/TheatreLM-v2.1-Characters" and "theatrelm_seed" not in row["source_mix"]:
errors.append("TheatreLM seeded rows must include theatrelm_seed in source_mix")
if row["source_dataset"] == "practical-dreamer/RPGPT_PublicDomain-alpaca" and "rpgpt_seed" not in row["source_mix"]:
errors.append("RPGPT seeded rows must include rpgpt_seed in source_mix")
if "transformation" in row and row["transformation"] != "seeded_synthetic_actor_json":
errors.append("transformation must be seeded_synthetic_actor_json when present")
if ("source_dataset" in row) != ("transformation" in row):
errors.append("source_dataset and transformation must appear together")
messages = row["messages"]
if not isinstance(messages, list) or len(messages) != 3:
errors.append("messages must contain exactly system, user, assistant messages")
return errors
expected_roles = ["system", "user", "assistant"]
for index, expected_role in enumerate(expected_roles):
message = messages[index]
if not isinstance(message, dict):
errors.append(f"message {index} must be an object")
continue
if set(message) != {"role", "content"}:
errors.append(f"message {index} must contain only role and content")
if message.get("role") != expected_role:
errors.append(f"message {index} role must be {expected_role}")
if not isinstance(message.get("content"), str) or not message["content"].strip():
errors.append(f"message {index} content must be non-empty text")
user_content = messages[1]["content"] if isinstance(messages[1], dict) else ""
for marker in ["premise:", "show_state JSON:", "actor JSON:", "director_instruction:"]:
if marker not in user_content:
errors.append(f"user message missing {marker}")
errors.extend(validate_assistant_content(messages[2]["content"], row, version, user_content))
return errors
def validate_assistant_content(content: str, row: dict[str, Any] | None = None, version: str = "v0", user_content: str = "") -> list[str]:
errors: list[str] = []
try:
value = json.loads(content)
except json.JSONDecodeError as exc:
return [f"assistant content JSON parse failed: {exc}"]
if not isinstance(value, dict):
return ["assistant content must parse to an object"]
keys = list(value)
if keys != OUTPUT_FIELDS:
errors.append(f"assistant keys must be exactly {OUTPUT_FIELDS}; got {keys}")
if version == "v1":
copied_fields = sorted(V1_BANNED_ASSISTANT_FIELDS & set(value))
if copied_fields:
errors.append(f"assistant must not copy state/input fields: {copied_fields}")
for field in OUTPUT_FIELDS:
if field not in value:
continue
if field == "tool_request":
continue
if field == "memory_update" and value[field] is None:
continue
if not isinstance(value[field], str):
errors.append(f"{field} must be a string")
continue
if not value[field].strip():
errors.append(f"{field} must not be empty")
line = value.get("line", "")
if isinstance(line, str):
words = line.split()
if not 6 <= len(words) <= 18:
errors.append(f"line must be 6-18 words; got {len(words)}")
lowered = line.lower()
banned = sorted(term for term in BANNED_TERMS if term in lowered)
if banned:
errors.append(f"line contains banned terms: {banned}")
intent = value.get("intent", "")
if isinstance(intent, str) and intent not in ALLOWED_INTENTS:
errors.append(f"intent must be one of {sorted(ALLOWED_INTENTS)}")
gesture = value.get("gesture", "")
if isinstance(gesture, str) and "_" in gesture:
errors.append("gesture must be a short theatrical phrase, not an enum-like token")
memory_update = value.get("memory_update", "")
if isinstance(memory_update, str):
if not memory_update.strip():
errors.append("memory_update must be null or non-empty text")
if len(memory_update) > 140:
errors.append("memory_update must be 140 characters or fewer")
elif memory_update is not None:
errors.append("memory_update must be null or a string")
if version == "v1" and row is not None:
errors.extend(validate_v1_finale_context(value, row, user_content))
tool_request = value.get("tool_request")
errors.extend(validate_tool_request(tool_request, version))
return errors
def validate_v1_finale_context(value: dict[str, Any], row: dict[str, Any], user_content: str) -> list[str]:
if value.get("intent") != "deliver_finale" and value.get("stage_effect") != "final_bow_lights":
return []
show_state = extract_show_state(user_content)
finale_requested = bool(show_state.get("finale_requested")) if isinstance(show_state, dict) else False
story_phase = show_state.get("story_phase") if isinstance(show_state, dict) else None
if row.get("row_type") == "finale" or finale_requested or story_phase == "finale":
return []
return ["deliver_finale/final_bow_lights only allowed for finale row, finale_requested, or story_phase finale"]
def extract_show_state(user_content: str) -> dict[str, Any] | None:
marker = "show_state JSON:"
next_marker = "\nactor JSON:"
if marker not in user_content:
return None
start = user_content.index(marker) + len(marker)
end = user_content.find(next_marker, start)
raw_json = user_content[start:end if end != -1 else None].strip()
try:
value = json.loads(raw_json)
except json.JSONDecodeError:
return None
return value if isinstance(value, dict) else None
def validate_tool_request(value: Any, version: str = "v0") -> list[str]:
if value is None:
return []
if not isinstance(value, dict):
return ["tool_request must be null or an object"]
if set(value) != {"tool", "args", "reason"}:
return ["tool_request must contain exactly tool, args, and reason"]
tool = value["tool"]
if tool not in ALLOWED_TOOLS:
return [f"tool_request tool must be one of {sorted(ALLOWED_TOOLS)}"]
if not isinstance(value["reason"], str) or not value["reason"].strip():
return ["tool_request reason must be non-empty text"]
if len(value["reason"]) > 140:
return ["tool_request reason must be 140 characters or fewer"]
args = value["args"]
if not isinstance(args, dict):
return ["tool_request args must be an object"]
tool_args = V1_TOOL_ARGS if version == "v1" else TOOL_ARGS
if set(args) != tool_args[tool]:
return [f"tool_request args for {tool} must be exactly {sorted(tool_args[tool])}"]
for arg_value in args.values():
if not isinstance(arg_value, str) or not arg_value.strip():
return ["tool_request arg values must be non-empty strings"]
if len(arg_value) > 120:
return ["tool_request arg values must be 120 characters or fewer"]
return []
def add_error(stats: dict[str, Any], max_errors: int, line_number: int, message: str) -> None:
stats["invalid"] += 1 if message.startswith("row JSON parse failed") else 0
if len(stats["errors"]) < max_errors:
stats["errors"].append(f"line {line_number}: {message}")
def count_lines(path: Path) -> int:
with path.open("r", encoding="utf-8") as handle:
return sum(1 for line in handle if line.strip())
def print_distribution(title: str, values: Counter) -> None:
print(f"{title}:")
for key, count in sorted(values.items()):
print(f" {key}: {count}")
if __name__ == "__main__":
main()