state-centric-learning-demo / code /encoding_generation /generate_all_domain_embeddings.py
vishalp's picture
Deploy State-Centric Learning live demo
dbc6675 verified
Raw
History Blame Contribute Delete
7.83 kB
"""
Generate embeddings for one all-domain tokenizer artifact.
This script creates one shared tokenizer artifact for the requested
domains, then embeds every split for every domain into a fixed space
compatible with a single pooled downstream model.
"""
from __future__ import annotations
import argparse
import os
import re
import sys
import numpy as np
import pddl
import pddl.logic.predicates
from tqdm import tqdm
from code.tokenization.multidomain import (
MultiDomainUnionTokenizer,
build_all_domain_tokenizer,
build_domain_specs,
save_tokenizer_manifest,
)
SPLITS = ["train", "validation", "test-interpolation", "test-extrapolation"]
_PREDICATE_REGEX = re.compile(r"\(([\w-]+(?: [\w-]+)*)\)")
def progress_enabled() -> bool:
return bool(sys.stdout.isatty())
def _extract_goal_atoms(problem) -> list[str]:
goals: list[str] = []
def visit(node):
if isinstance(node, pddl.logic.predicates.Predicate):
args = [t.name if hasattr(t, "name") else str(t) for t in node.terms]
goals.append(f"({node.name} {' '.join(args)})")
elif hasattr(node, "operands"):
for op in node.operands:
visit(op)
elif hasattr(node, "_operands"):
for op in node._operands:
visit(op)
visit(problem.goal)
return goals
def _get_objects(problem, domain) -> list[str]:
objs = {o.name for o in problem.objects}
objs.update(o.name for o in domain.constants)
return sorted(objs)
def main() -> None:
parser = argparse.ArgumentParser(
description="Generate embeddings with one all-domain tokenizer artifact."
)
parser.add_argument(
"--tokenizer",
required=True,
choices=["wl", "simhash", "shortest_path", "graphbpe", "random"],
)
parser.add_argument(
"--domains",
nargs="+",
required=True,
help="Domains included in the shared tokenizer fit.",
)
parser.add_argument("--data_dir", default="data")
parser.add_argument("--output_dir", required=True)
parser.add_argument("--model_dir", required=True)
parser.add_argument(
"--strategy",
choices=["auto", "pooled", "union"],
default="auto",
help="All-domain tokenizer fit strategy.",
)
parser.add_argument("--iterations", type=int, default=2, help="WL iterations")
parser.add_argument("--hash_dim", type=int, default=128, help="SimHash dimension")
parser.add_argument("--seed", type=int, default=42, help="Tokenizer seed")
parser.add_argument("--max_path_length", type=int, default=5)
parser.add_argument("--vocab_size", type=int, default=1000)
parser.add_argument("--num_iterations", type=int, default=100)
parser.add_argument("--random_dim", type=int, default=128)
parser.add_argument(
"--no_random_normalize",
action="store_true",
help="Disable unit normalization for the random baseline",
)
args = parser.parse_args()
os.makedirs(args.output_dir, exist_ok=True)
os.makedirs(args.model_dir, exist_ok=True)
tokenizer_kwargs = {
"iterations": args.iterations,
"hash_dim": args.hash_dim,
"seed": args.seed,
"max_path_length": args.max_path_length,
"vocab_size": args.vocab_size,
"num_iterations": args.num_iterations,
"random_dim": args.random_dim,
"normalize": (not args.no_random_normalize),
}
domain_specs = build_domain_specs(args.data_dir, args.domains)
tokenizer, resolved_strategy = build_all_domain_tokenizer(
args.tokenizer,
domain_specs,
strategy=args.strategy,
**tokenizer_kwargs,
)
manifest_path = os.path.join(
args.model_dir,
f"all_domains_{args.tokenizer}.json",
)
save_tokenizer_manifest(
tokenizer=tokenizer,
manifest_path=manifest_path,
tokenizer_name=args.tokenizer,
domains=args.domains,
fit_strategy=resolved_strategy,
tokenizer_kwargs=tokenizer_kwargs,
)
print(f"Saved all-domain tokenizer manifest to {manifest_path}")
print(
f"Resolved strategy: {resolved_strategy} | "
f"Embedding dimension: {tokenizer.get_embedding_dim()}"
)
for domain_name in args.domains:
domain_pddl = os.path.join(args.data_dir, "pddl", domain_name, "domain.pddl")
domain = pddl.parse_domain(domain_pddl)
if isinstance(tokenizer, MultiDomainUnionTokenizer):
tokenizer.set_active_domain(domain_name, domain_pddl)
print(f"\n{'=' * 60}")
print(f"Embedding domain {domain_name} with shared tokenizer [{args.tokenizer}]")
print(f"{'=' * 60}")
for split in SPLITS:
split_state_dir = os.path.join(args.data_dir, "states", domain_name, split)
split_pddl_dir = os.path.join(args.data_dir, "pddl", domain_name, split)
split_out_dir = os.path.join(args.output_dir, domain_name, split)
os.makedirs(split_out_dir, exist_ok=True)
if not os.path.exists(split_state_dir):
print(f"Skipping {domain_name}/{split}: state dir not found")
continue
traj_files = sorted(
file_name
for file_name in os.listdir(split_state_dir)
if file_name.endswith(".traj")
)
for traj_file in tqdm(
traj_files,
desc=f"{domain_name}/{split}",
disable=(not progress_enabled()),
):
prob_name = traj_file.replace(".traj", "")
prob_pddl = os.path.join(split_pddl_dir, f"{prob_name}.pddl")
traj_path = os.path.join(split_state_dir, traj_file)
out_traj_path = os.path.join(split_out_dir, f"{prob_name}.npy")
out_goal_path = os.path.join(split_out_dir, f"{prob_name}_goal.npy")
if not os.path.exists(prob_pddl):
continue
try:
problem = pddl.parse_problem(prob_pddl)
objects = _get_objects(problem, domain)
goal_atoms = _extract_goal_atoms(problem)
with open(traj_path, "r", encoding="utf-8") as f:
lines = f.readlines()
state_embeddings = []
for line in lines:
state_atoms = [f"({a})" for a in _PREDICATE_REGEX.findall(line.strip())]
try:
emb = tokenizer.transform_state(
state_atoms,
goal_atoms,
objects,
problem_pddl_path=prob_pddl,
)
except TypeError:
emb = tokenizer.transform_state(state_atoms, goal_atoms, objects)
state_embeddings.append(emb)
traj_matrix = np.array(state_embeddings, dtype=np.float32)
try:
goal_vec = tokenizer.transform_goal(
goal_atoms,
objects,
problem_pddl_path=prob_pddl,
)
except TypeError:
goal_vec = tokenizer.transform_goal(goal_atoms, objects)
goal_vec = goal_vec.astype(np.float32)
np.save(out_traj_path, traj_matrix)
np.save(out_goal_path, goal_vec)
except Exception as exc:
print(f"Error embedding {domain_name}/{split}/{prob_name}: {exc}")
print("\nDone.")
if __name__ == "__main__":
main()