Yui-home-assistant / train /taskgen.py
ABarroso647
Ship Yui as a ZeroGPU Gradio Space: two-stage voice assistant
9275c01 unverified
Raw
History Blame Contribute Delete
23.1 kB
"""Grounded task generator for Brain 1: one base generator, two output formats.
gen: propose (home, utterance, gold call) -> execute in SandboxHome -> keep only
verified records (actions must produce the expected diff; rejections must fail
with a NAME miss). Records land in data/processed/tasks.jsonl, self-contained.
sft: tasks.jsonl -> intent_sft-shaped rows (Brain 1 stage-2 SFT)
rl: tasks.jsonl -> prompt + expected-diff rows (RLVR; sandbox replays = reward)
Run: uv run --group sandbox --group train python train/taskgen.py gen|sft|rl
"""
import json
import random
import sys
from collections import Counter
from dataclasses import asdict, dataclass
from datetime import datetime, timedelta
from pathlib import Path
import yaml
ROOT = Path(__file__).resolve().parent.parent
sys.path.insert(0, str(ROOT))
from sandbox.harness import SandboxHome # noqa: E402
_CFG = yaml.safe_load((Path(__file__).parent / "config.yaml").read_text())
PROCESSED = ROOT / "data" / "processed"
TASKS_FILE = PROCESSED / "tasks.jsonl"
HOME_DIRS = [
ROOT / "data" / "homes",
ROOT / "plans" / "home-assistant-datasets-src" / "datasets" / "assist-mini",
]
OFF_LIKE = {"off", "closed", "unlocked", "docked", "idle"}
TEMPLATES = {
("light", "on"): ["turn on the {name}", "switch the {name} on", "can you turn the {name} on", "{name} on please"],
("light", "off"): ["turn off the {name}", "switch the {name} off", "kill the {name}", "{name} off please"],
("cover", "on"): ["open the {name}", "can you open the {name}", "pull up the {name}"],
("cover", "off"): ["close the {name}", "shut the {name}", "can you close the {name}"],
("lock", "on"): ["lock the {name}", "secure the {name}", "make sure the {name} is locked"],
("lock", "off"): ["unlock the {name}", "open up the {name}", "can you unlock the {name}"],
("valve", "on"): ["open the {name}", "turn on the {name}"],
("valve", "off"): ["close the {name}", "shut off the {name}"],
(None, "on"): ["turn on the {name}", "switch on the {name}", "start the {name}"],
(None, "off"): ["turn off the {name}", "switch off the {name}", "shut down the {name}"],
}
ATTR_TEMPLATES = {
"brightness": ["set the {name} to {pct} percent", "dim the {name} to {pct}%", "make the {name} {pct} percent bright"],
"position": ["set the {name} to {pct} percent open", "open the {name} to {pct}%"],
}
MEDIA = {
"playing": ("HassMediaPause", ["pause the {name}", "pause whatever is on the {name}"]),
"paused": ("HassMediaUnpause", ["resume the {name}", "unpause the {name}"]),
}
# automation requests are NOT templated here: subagents write them per home from
# train/prompts/routegen.md into ROUTES_DIR; gen() ingests + validates them.
ROUTES_DIR = PROCESSED / "route_requests"
BATCH_TEMPLATES = {
"on": ["turn on all the lights", "switch every light on", "all lights on please"],
"off": ["turn off all the lights", "switch every light off", "kill all the lights"],
}
EXCLUSION_TEMPLATES = [
"turn off all the lights except the {kept}",
"switch off every light but the {kept}",
"turn everything off except the {kept}",
]
DENIALS = [
"Sorry, I can't find a {name} to control.",
"I don't see a {name} in this home.",
"There's no {name} listed here.",
]
FOREIGN = [
("hot tub heater", "switch"), ("pool light", "light"), ("wine fridge", "switch"),
("doggy door", "cover"), ("side gate lock", "lock"), ("attic fan", "fan"),
("sprinkler valve", "valve"), ("sauna", "climate"), ("garage heater", "switch"),
("greenhouse light", "light"),
]
CONTROLLABLE = {"light", "switch", "fan", "cover", "lock", "valve", "media_player"}
@dataclass
class TaskRecord:
home: str
category: str # control | attr | disambig | reject_missing
utterance: str
devices_block: str
time: str
# list of {"name": intent, "slots": {...}} (batch => several); None => no call
# (rejections AND route_automation — those route verbatim, see to_sft)
gold: list | None
# intent-relevant outcomes only: {entity: {"state"|attr: after_value}}.
# Reward = containment (replay diff includes these after-values), NOT exact
# diff equality — incidental keys (color_mode, hidden brightness restores)
# aren't reproducible from yaml and aren't what the record asserts.
expect: dict
response: str | None # denial sentence for rejections
inventory: str # raw home yaml so RL reward can rebuild the sandbox
meta: dict | None = None # route records: {"trigger": eid, "actuator": eid} for the Brain-2 reward
def _rand_time(rng: random.Random) -> str:
dt = datetime(2024, 1, 1) + timedelta(minutes=rng.randrange(2 * 365 * 24 * 60))
hour = dt.strftime("%I").lstrip("0")
return f"{hour}:{dt.minute} {dt.strftime('%p')} on {dt.strftime('%A, %B')} {dt.day} {dt.year}"
# dup-named entities in verb-ambiguous domains get the domain noun appended in the
# UTTERANCE (the slot keeps the real name) so the user text justifies the domain slot
_DOMAIN_NOUN = {"light": "light", "fan": "fan", "switch": "switch"}
def _propose(entity: dict, state: str, attrs: dict, dup: bool, rng: random.Random):
"""Pick an intent that should change this entity, given its current state."""
domain, name = entity["id"].split(".")[0], entity.get("name", entity["id"])
slots = {"name": name} | ({"domain": domain} if dup else {})
spoken = name
if dup:
if domain in _DOMAIN_NOUN:
spoken = f"{name} {_DOMAIN_NOUN[domain]}"
elif domain not in ("cover", "lock", "valve"):
return None # no way to phrase the disambiguation unambiguously
if domain == "media_player" and state in MEDIA:
intent, templates = MEDIA[state]
return "control", intent, slots, rng.choice(templates).format(name=spoken)
if domain == "light" and state == "on" and rng.random() < 0.5:
current = round((attrs.get("brightness") or 0) / 255 * 100)
if choices := [p for p in (10, 25, 40, 60, 75, 90) if abs(p - current) >= 15]:
pct = rng.choice(choices)
slots["brightness"] = pct
return "attr", "HassLightSet", slots, rng.choice(ATTR_TEMPLATES["brightness"]).format(name=spoken, pct=pct)
if domain == "cover" and rng.random() < 0.4:
current = attrs.get("current_position") or 0
if choices := [p for p in (20, 35, 50, 65, 80) if abs(p - current) >= 15]:
pct = rng.choice(choices)
slots["position"] = pct
return "attr", "HassSetPosition", slots, rng.choice(ATTR_TEMPLATES["position"]).format(name=spoken, pct=pct)
direction = "on" if state in OFF_LIKE else "off"
intent = {"on": "HassTurnOn", "off": "HassTurnOff"}[direction]
templates = TEMPLATES.get((domain, direction)) or TEMPLATES[(None, direction)]
return "control", intent, slots, rng.choice(templates).format(name=spoken)
# every attr devices_block renders must be patched back, or replayed homes
# won't reproduce the prompt's device lines
_PATCH_ATTRS = ("brightness", "current_position", "percentage", "volume_level", "media_title")
def _live_inventory(home: SandboxHome) -> str:
"""Inventory yaml with current state written back, so a fresh boot reproduces
exactly the world the record's devices_block describes (drift-consistent)."""
snap = home.snapshot()
inv = yaml.safe_load(yaml.safe_dump(home.inventory))
for e in inv.get("entities", []):
state, attrs = snap[e["id"]]
e["state"] = state
for k in _PATCH_ATTRS:
if attrs.get(k) is not None:
e.setdefault("attributes", {})[k] = attrs[k]
else:
e.get("attributes", {}).pop(k, None) # stale value would resurrect on boot
return yaml.safe_dump(inv, sort_keys=False)
def gen_home(path: Path, n: int, rng: random.Random) -> list[TaskRecord]:
home = SandboxHome.from_file(path)
home_id = path.parent.name if path.name == "_fixtures.yaml" else path.stem
entities = [e for e in home.inventory.get("entities", []) if e["id"].split(".")[0] in CONTROLLABLE]
name_counts = Counter(e.get("name", e["id"]) for e in home.inventory.get("entities", []))
records: list[TaskRecord] = []
attempts = 0
while len(records) < n and attempts < n * 4 and entities:
attempts += 1
block, time = home.devices_block(), _rand_time(rng)
inv_yaml = _live_inventory(home)
if rng.random() < 0.2: # rejection: a plausible device this home doesn't have
name, domain = rng.choice([f for f in FOREIGN if f[0] not in name_counts])
templates = TEMPLATES.get((domain, "on")) or TEMPLATES[(None, "on")]
probe = home.execute("HassTurnOn", {"name": name})
if probe.ok or "NAME" not in (probe.error or ""):
continue # grounding: only keep if HA truly can't match it
records.append(TaskRecord(
home_id, "reject_missing", rng.choice(templates).format(name=name),
block, time, None, {}, rng.choice(DENIALS).format(name=name), inv_yaml,
))
continue
entity = rng.choice(entities)
state, attrs = home.snapshot()[entity["id"]]
dup = name_counts[entity.get("name", entity["id"])] > 1
proposal = _propose(entity, state, attrs, dup, rng)
if proposal is None:
continue
category, intent, slots, utterance = proposal
result = home.execute(intent, slots)
if not result.ok or entity["id"] not in result.changed:
continue # grounding: drop anything that didn't verifiably act
delta = result.changed[entity["id"]]
want = {}
if "state" in delta:
want["state"] = delta["state"][1]
for slot_key, diff_key in (("brightness", "brightness"), ("position", "current_position")):
if slot_key in slots and diff_key in delta:
want[diff_key] = delta[diff_key][1]
if not want:
continue
records.append(TaskRecord(
home_id, "disambig" if dup else category, utterance, block, time,
[{"name": intent, "slots": slots}], {entity["id"]: want}, None, inv_yaml,
))
records += _gen_routes(home, home_id, rng)
records += _gen_batches(home, home_id, name_counts, rng)
home.close()
return records
def _gen_routes(home: SandboxHome, home_id: str, rng: random.Random) -> list[TaskRecord]:
"""Ingest subagent-written automation requests (ROUTES_DIR/<home>.json),
validating each against the live inventory. Doubles as the Brain-2 GRPO
prompt set; meta carries trigger/actuator for its reward."""
path = ROUTES_DIR / f"{home_id}.json"
if not path.exists():
return []
ents = home.inventory.get("entities", [])
ids = {e["id"] for e in ents}
name_of = {e["id"]: e.get("name", e["id"]) for e in ents}
records, seen = [], set()
for req in json.loads(path.read_text()):
utt = (req.get("utterance") or "").strip()
trig, act = req.get("trigger"), req.get("actuator")
problem = (
"empty" if not utt
else "dup" if utt.lower() in seen
else "bad actuator" if act not in ids
else "bad trigger" if trig not in ids and trig not in ("time", "sun")
else "actuator name missing from utterance" if name_of[act].lower() not in utt.lower()
else None
)
if problem:
print(f"[gen] drop route ({home_id}, {problem}): {utt[:60]!r}")
continue
seen.add(utt.lower())
records.append(TaskRecord(
home_id, "route_automation", utt, home.devices_block(), _rand_time(rng),
None, {}, None, _live_inventory(home),
meta={"trigger": trig, "actuator": act},
))
return records
def _gen_batches(home: SandboxHome, home_id: str, name_counts: Counter, rng: random.Random) -> list[TaskRecord]:
"""One batch + one exclusion record per home with enough lights. Gold lists every
targeted light; grounded by executing all calls and requiring the changers to change."""
records = []
lights = [e for e in home.inventory.get("entities", []) if e["id"].startswith("light.")]
def call(e, intent):
slots = {"name": e.get("name", e["id"])}
if name_counts[e.get("name", e["id"])] > 1:
slots["domain"] = "light"
return {"name": intent, "slots": slots}
def run_all(targets, intent, movers):
block, time, inv = home.devices_block(), _rand_time(rng), _live_inventory(home)
gold = [call(e, intent) for e in targets]
expect = {}
for e in movers:
c = call(e, intent)
r = home.execute(c["name"], c["slots"])
if not r.ok or e["id"] not in r.changed or "state" not in r.changed[e["id"]]:
return None
expect[e["id"]] = {"state": r.changed[e["id"]]["state"][1]}
return block, time, inv, gold, expect
live = home.states()
on = [e for e in lights if live[e["id"]] == "on"]
off = [e for e in lights if live[e["id"]] == "off"]
direction, movers = ("off", on) if len(on) >= len(off) else ("on", off)
if len(movers) >= 2:
done = run_all(lights, "HassTurnOff" if direction == "off" else "HassTurnOn", movers)
if done:
block, time, inv, gold, expect = done
records.append(TaskRecord(
home_id, "batch", rng.choice(BATCH_TEMPLATES[direction]),
block, time, gold, expect, None, inv,
))
live = home.states()
on = [e for e in lights if live[e["id"]] == "on"]
if len(on) >= 3:
kept = rng.choice(on)
targets = [e for e in on if e["id"] != kept["id"]]
done = run_all(targets, "HassTurnOff", targets)
if done:
block, time, inv, gold, expect = done
records.append(TaskRecord(
home_id, "exclusion",
rng.choice(EXCLUSION_TEMPLATES).format(kept=kept.get("name", kept["id"])),
block, time, gold, expect, None, inv,
))
return records
def _find_homes(roots: list[Path] = HOME_DIRS) -> list[Path]:
files: list[Path] = []
for root in roots:
if root.exists():
files += sorted(root.glob("**/_fixtures.yaml")) + sorted(root.glob("*.yaml"))
return files
def gen(n_per_home: int = 8, seed: int = 0) -> None:
PROCESSED.mkdir(parents=True, exist_ok=True)
rng = random.Random(seed)
records: list[TaskRecord] = []
# publishable output: ONLY our generated pool, never third-party fixtures
for path in _find_homes([HOME_DIRS[0]]):
recs = gen_home(path, n_per_home, rng)
records += recs
print(f"[gen] {path.parent.name}: {len(recs)} records")
with TASKS_FILE.open("w") as f:
for r in records:
f.write(json.dumps(asdict(r), ensure_ascii=False) + "\n")
print(f"[gen] {len(records)} total {dict(Counter(r.category for r in records))} -> {TASKS_FILE}")
def _system(record: dict) -> str:
return (
f"{_CFG['system_prompt']}\n\nActions:\n{_CFG['intents']}"
f"\n\nDevices:\n{record['devices_block']}\nCurrent time: {record['time']}"
)
def to_sft(record: dict) -> dict:
if record["category"] == "route_automation":
payload = {"automation": record["utterance"]} # verbatim handoff to Brain 2
elif record["gold"]:
payload = {"intents": record["gold"]}
else:
payload = {"response": record["response"]}
return {
"messages": [
{"role": "system", "content": _system(record)},
{"role": "user", "content": record["utterance"]},
{"role": "assistant", "content": json.dumps(payload, ensure_ascii=False, separators=(",", ":"))},
],
"category": record["category"],
"home": record["home"],
}
def to_rl(record: dict) -> dict:
meta = record.get("meta") or {}
return {
"messages": [
{"role": "system", "content": _system(record)},
{"role": "user", "content": record["utterance"]},
],
"category": record["category"],
"home": record["home"],
"inventory": record["inventory"],
"expect": json.dumps(record["expect"], ensure_ascii=False),
"gold": json.dumps(record["gold"], ensure_ascii=False),
"utterance": record["utterance"],
"devices_block": record["devices_block"],
"trigger": meta.get("trigger", ""),
"actuator": meta.get("actuator", ""),
}
def _format(kind: str) -> None:
from datasets import Dataset
fmt = {"sft": to_sft, "rl": to_rl}[kind]
rows = [fmt(json.loads(line)) for line in TASKS_FILE.read_text().splitlines()]
out = PROCESSED / "taskgen" / f"{kind}.parquet"
out.parent.mkdir(exist_ok=True)
Dataset.from_list(rows).to_parquet(str(out))
print(f"[{kind}] {len(rows)} rows -> {out}")
def expect_met(changed: dict, expect: dict) -> bool:
"""The RLVR reward core: every expected after-value appears in the diff."""
return all(
eid in changed and key in changed[eid] and changed[eid][key][1] == after
for eid, wants in expect.items()
for key, after in wants.items()
)
def audit(n: int = 12, seed: int = 11) -> None:
"""Replay n random action records on fresh boots; block + expect must reproduce."""
lines = [json.loads(line) for line in TASKS_FILE.read_text().splitlines()]
sample = random.Random(seed).sample([r for r in lines if r["gold"]], n)
ok = 0
for rec in sample:
home = SandboxHome(rec["inventory"])
block_ok = rec["devices_block"] == home.devices_block()
merged: dict = {}
calls_ok = True
for call in rec["gold"]:
result = home.execute(call["name"], call["slots"])
calls_ok = calls_ok and result.ok
for eid, delta in result.changed.items():
merged.setdefault(eid, {}).update(delta)
met = calls_ok and expect_met(merged, rec["expect"])
home.close()
ok += block_ok and met
if not (block_ok and met):
print(f"[audit] MISMATCH {rec['home']} {rec['utterance']!r} block={block_ok} expect={met}")
print(f"[audit] {ok}/{n} records replay perfectly")
def diversity() -> None:
"""Score route_automation utterances on diversity; flag template-y sameness.
Checks: near-duplicate pairs (Jaccard>0.7), opener concentration, type-token
ratio, trigger-kind spread. Run after ingesting a new generation batch."""
lines = [json.loads(line) for line in TASKS_FILE.read_text().splitlines()]
routes = [r for r in lines if r["category"] == "route_automation"]
if not routes:
print("[diversity] no route_automation records")
return
utts = [r["utterance"].lower().replace(",", "") for r in routes]
word_sets = [set(u.split()) for u in utts]
dups = [
(utts[i], utts[j])
for i in range(len(utts)) for j in range(i + 1, len(utts))
if len(word_sets[i] & word_sets[j]) / len(word_sets[i] | word_sets[j]) > 0.7
]
trigrams = [" ".join(u.split()[i:i + 3]) for u in utts for i in range(len(u.split()) - 2)]
distinct_tri = len(set(trigrams)) / len(trigrams)
openers = Counter(" ".join(u.split()[:2]) for u in utts)
kinds = Counter(
"time/sun" if r["meta"]["trigger"] in ("time", "sun") else r["meta"]["trigger"].split(".")[0]
for r in routes
)
top_opener, top_n = openers.most_common(1)[0]
top_kind, kind_n = kinds.most_common(1)[0]
print(f"[diversity] {len(routes)} routes over {len({r['home'] for r in routes})} homes")
print(f"[diversity] trigger kinds: {dict(kinds)} — top {top_kind} {kind_n / len(routes):.0%} (flag > 50%)")
print(f"[diversity] distinct trigrams: {distinct_tri:.2f} (flag < 0.55)")
print(f"[diversity] top opener: {top_opener!r} x{top_n} ({top_n / len(utts):.0%}, flag > 25%)")
print(f"[diversity] near-duplicate pairs (jaccard>0.7): {len(dups)} (flag > {len(utts) // 20})")
for a, b in dups[:5]:
print(f" ~ {a!r}\n {b!r}")
flags = (
(distinct_tri < 0.55) + (top_n / len(utts) > 0.25)
+ (len(dups) > len(utts) // 20) + (kind_n / len(routes) > 0.5)
)
print(f"[diversity] {'FAIL — regenerate the flagged areas' if flags else 'PASS'}")
def show(per_category: int = 2, seed: int = 3) -> None:
"""Pretty-print sample records per category, plus one full SFT and RL row."""
lines = [json.loads(line) for line in TASKS_FILE.read_text().splitlines()]
rng = random.Random(seed)
by_cat: dict[str, list] = {}
for r in lines:
by_cat.setdefault(r["category"], []).append(r)
for cat, recs in sorted(by_cat.items()):
print(f"\n=== {cat} ({len(recs)} records) ===")
for rec in rng.sample(recs, min(per_category, len(recs))):
print(f" [{rec['home']}] user: {rec['utterance']!r}")
if rec["gold"]:
print(f" gold: {json.dumps(rec['gold'], separators=(',', ':'))}")
print(f" expect: {json.dumps(rec['expect'], separators=(',', ':'))}")
elif rec["category"] == "route_automation":
print(f" gold: route -> Brain 2, verbatim (meta {rec['meta']})")
else:
print(f" gold: (no tool call) -> {rec['response']!r}")
rec = rng.choice([r for r in lines if r["gold"]])
sft_row = to_sft(rec)
print("\n=== one full SFT row (what Brain 1 trains on) ===")
for m in sft_row["messages"]:
print(f"--- {m['role']} ---\n{m['content']}")
rl_row = to_rl(rec)
print("\n=== same record as RL row (no assistant turn; sandbox judges) ===")
print("messages: system + user as above, NO assistant")
print(f"expect: {rl_row['expect']}")
print(f"inventory: {len(rl_row['inventory'])} chars of patched home yaml")
def check() -> None:
"""Boot-validate every home in the pool; report stats or the failure."""
ok = bad = 0
for path in _find_homes():
label = path.parent.name if path.name == "_fixtures.yaml" else path.stem
try:
home = SandboxHome.from_file(path)
except Exception as err: # noqa: BLE001
bad += 1
print(f"[check] FAIL {label}: {str(err)[:120]}")
continue
ents = home.inventory.get("entities", [])
domains = Counter(e["id"].split(".")[0] for e in ents)
dups = [n for n, c in Counter(e.get("name") for e in ents).items() if c > 1]
home.close()
ok += 1
print(f"[check] ok {label}: {len(ents)} entities {dict(domains)}"
+ (f" dups={dups}" if dups else ""))
print(f"[check] {ok} ok, {bad} failed")
if __name__ == "__main__":
cmd = sys.argv[1] if len(sys.argv) > 1 else "gen"
{"gen": gen, "sft": lambda: _format("sft"), "rl": lambda: _format("rl"), "check": check,
"audit": audit, "show": show, "diversity": diversity}[cmd]()