Download src/dooable/cli.py from ChatterjeeLab/DooABLe: direct link, hf CLI and curl.
- Browser
- Download file 6.71 kB
-
https://huggingface.co/ChatterjeeLab/DooABLe/resolve/main/src/dooable/cli.py
- Command line
-
hf download hf://ChatterjeeLab/DooABLe/src/dooable/cli.py
-
curl -L -o cli.py https://huggingface.co/ChatterjeeLab/DooABLe/resolve/main/src/dooable/cli.py
6.71 kB
| """Command-line entry points for data, training, generation, and verification.""" | |
| import argparse, json | |
| from pathlib import Path | |
| import numpy as np | |
| from .graph import Graph, toy_graph, grid_graph | |
| from .exact import ( | |
| solve, | |
| sample, | |
| endpoint_distribution, | |
| expected_cost, | |
| uniform_policy, | |
| tilted_reference, | |
| ) | |
| def main(): | |
| p = argparse.ArgumentParser(prog="dooable") | |
| sub = p.add_subparsers(dest="command", required=True) | |
| b = sub.add_parser("build") | |
| b.add_argument("--parents", required=True) | |
| b.add_argument("--reagents", required=True) | |
| b.add_argument("--reactions", required=True) | |
| b.add_argument("--budget", type=int, default=2) | |
| b.add_argument("--max-nodes", type=int, default=100000) | |
| b.add_argument("--parent-limit", type=int) | |
| b.add_argument("--output", required=True) | |
| t = sub.add_parser("toy") | |
| t.add_argument("--kind", choices=["multiplicity", "grid"], default="multiplicity") | |
| t.add_argument("--output", required=True) | |
| t.add_argument("--budget", type=int, default=6) | |
| for name in ["exact", "train"]: | |
| a = sub.add_parser(name) | |
| a.add_argument("--graph", required=True) | |
| a.add_argument("--rewards") | |
| a.add_argument("--temperature", type=float, default=1.0) | |
| a.add_argument("--output", required=True) | |
| a.add_argument("--seed", type=int, default=0) | |
| if name == "train": | |
| a.add_argument("--steps", type=int, default=2000) | |
| a.add_argument("--batch-size", type=int, default=64) | |
| a.add_argument( | |
| "--backward", | |
| choices=["learned", "uniform", "exact", "unnormalized"], | |
| default="learned", | |
| ) | |
| a.add_argument("--resume") | |
| a.add_argument("--initialize") | |
| a = sub.add_parser("sample") | |
| a.add_argument("--model", required=True) | |
| a.add_argument("--n", type=int, default=1000) | |
| a.add_argument("--seed", type=int, default=0) | |
| a.add_argument("--output", required=True) | |
| a = sub.add_parser("fit-properties") | |
| a.add_argument("--data", default="data/downloads") | |
| a.add_argument("--output", required=True) | |
| a.add_argument("--seed", type=int, default=0) | |
| a = sub.add_parser("score") | |
| a.add_argument("--graph", required=True) | |
| a.add_argument("--models") | |
| a.add_argument("--output", required=True) | |
| a = sub.add_parser("replay") | |
| a.add_argument("--samples", required=True) | |
| a.add_argument("--budget", type=int, required=True) | |
| args = p.parse_args() | |
| if args.command == "build": | |
| from .chemistry import read_catalog, reactions, build_graph | |
| g = build_graph( | |
| read_catalog(args.parents, args.parent_limit), | |
| read_catalog(args.reagents), | |
| reactions(args.reactions), | |
| args.budget, | |
| args.max_nodes, | |
| ) | |
| g.save(args.output) | |
| print( | |
| json.dumps( | |
| { | |
| "nodes": len(g.nodes), | |
| "edges": len(g.edges), | |
| "outcomes": len(g.terminals), | |
| } | |
| ) | |
| ) | |
| return | |
| if args.command == "toy": | |
| g = toy_graph() if args.kind == "multiplicity" else grid_graph(6, args.budget) | |
| g.save(args.output) | |
| return | |
| if args.command in ["exact", "train", "score"]: | |
| g = Graph.load(args.graph) | |
| if args.command == "score": | |
| Path(args.output).parent.mkdir(parents=True, exist_ok=True) | |
| if args.models: | |
| from .properties import property_rewards | |
| rewards, scores = property_rewards(g, args.models) | |
| scores.to_csv(Path(args.output).with_suffix(".csv"), index=False) | |
| else: | |
| from .chemistry import descriptor_rewards | |
| rewards = descriptor_rewards(g) | |
| Path(args.output).write_text(json.dumps(rewards, indent=2)) | |
| return | |
| rewards = ( | |
| json.loads(Path(args.rewards).read_text()) | |
| if args.rewards | |
| else {y: 0.0 for y in g.terminals} | |
| ) | |
| out = Path(args.output) | |
| out.mkdir(parents=True, exist_ok=True) | |
| if args.command == "exact": | |
| sol = solve(g, rewards, args.temperature) | |
| g.save(out / "graph.json") | |
| np.save(out / "forward.npy", sol.forward) | |
| (out / "solution.json").write_text( | |
| json.dumps( | |
| { | |
| "temperature": args.temperature, | |
| "target": sol.target, | |
| "log_z": sol.log_z, | |
| "expected_cost": expected_cost(g, sol.forward), | |
| }, | |
| indent=2, | |
| ) | |
| ) | |
| else: | |
| from .learning import train | |
| _, h = train( | |
| g, | |
| rewards, | |
| args.temperature, | |
| args.steps, | |
| args.batch_size, | |
| seed=args.seed, | |
| output=out, | |
| backward=args.backward, | |
| resume=args.resume, | |
| initialize=args.initialize, | |
| ) | |
| print(json.dumps(h[-1])) | |
| return | |
| if args.command == "sample": | |
| d = Path(args.model) | |
| if (d / "forward.npy").exists(): | |
| g = Graph.load(d / "graph.json") | |
| forward = np.load(d / "forward.npy") | |
| else: | |
| from .learning import load_model | |
| model, _ = load_model(d) | |
| g = model.graph | |
| forward = model.probabilities() | |
| rows = sample(g, forward, args.n, args.seed) | |
| Path(args.output).parent.mkdir(parents=True, exist_ok=True) | |
| Path(args.output).write_text("".join(json.dumps(r) + "\n" for r in rows)) | |
| return | |
| if args.command == "fit-properties": | |
| from .properties import fit_property | |
| print( | |
| json.dumps( | |
| [ | |
| fit_property(n, args.data, args.output, args.seed) | |
| for n in ["caco2", "bace"] | |
| ], | |
| indent=2, | |
| ) | |
| ) | |
| return | |
| if args.command == "replay": | |
| from .chemistry import replay | |
| rows = [json.loads(l) for l in Path(args.samples).read_text().splitlines() if l] | |
| passed = sum(replay(r, args.budget) for r in rows) | |
| print( | |
| json.dumps( | |
| { | |
| "paths": len(rows), | |
| "replayed": passed, | |
| "fraction": passed / max(len(rows), 1), | |
| } | |
| ) | |
| ) | |
| if passed != len(rows): | |
| raise SystemExit(1) | |
| if __name__ == "__main__": | |
| main() | |