AutoParallel provides three serialization APIs for different workflows:
| API | What it saves | Size | Use case |
|---|---|---|---|
save() / load() |
Full optimizer: graph, strategies, costs, constraints, solution | Large (~300MB for LLaMA-3 8B) | Offline exploration, re-solving, what-if analysis |
save_placements() / load_placements() |
Per-node placement choices (output + input specs) | Small (~100KB) | Reapplying a solution in a training script |
get_json() |
Rich export for visualization (nodes, edges, costs, clusters, source info) | Medium (~2MB) | Feeding the HTML visualizer |
Run the expensive tracing + optimization once, save the full state, then explore interactively without the model code or a process group.
# === Script: trace and save ===
with AutoParallel(model, input_fn, mesh, mp_policy) as autop:
autop.add_input_constraints([x_sharding])
autop.add_output_constraints([out_sharding])
autop.optimize_placement()
autop.sharding_optimizer.save("model.ap")# === Notebook: load and explore ===
from autoparallel.optimize_sharding import ShardingOptimizer
opt = ShardingOptimizer.load("model.ap")
# Visualize
from autoparallel.visualizer.build_display_from_json import generate_visualization_html
from IPython.display import HTML
HTML(generate_visualization_html(opt.get_json()))
# Inspect a node
opt.print_costs_for_node(opt.nodes[42])
# What-if: constrain a node and re-solve
from torch.distributed.tensor.placement_types import Shard, Replicate
original = opt.get_solution()
names = opt.add_node_constraint(node, (Shard(0), Replicate()))
new = opt.resolve()
opt.diff_solutions(original, new)
# Revert
opt.remove_constraints(names)
opt.resolve()No original model code or distributed process group is needed in the notebook. A GPU is not required for loading, inspecting, or re-solving.
Save the optimizer's placement choices as a lightweight JSON file, then reuse them in a later run without calling optimize_placement() again. Tracing and ILP construction still happen — only the solve step is skipped.
# === First run: solve and save placements ===
with AutoParallel(model, input_fn, mesh, mp_policy) as autop:
autop.add_input_constraints([x_sharding])
autop.add_output_constraints([out_sharding])
solution = autop.optimize_placement()
autop.sharding_optimizer.save_placements("placements.json")
module = autop.apply_placement(solution)# === Later run: load placements instead of re-solving ===
with AutoParallel(model, input_fn, mesh, mp_policy) as autop:
autop.add_input_constraints([x_sharding])
autop.add_output_constraints([out_sharding])
solution = autop.sharding_optimizer.load_placements("placements.json")
module = autop.apply_placement(solution)The placements file is a small JSON with the mesh shape, dim names, and per-node output + input placement strings. load_placements validates that the mesh matches and finds the exact strategy by matching both output and input specs.
with AutoParallel(model, input_fn, mesh, mp_policy) as autop:
autop.add_input_constraints([x_sharding])
autop.add_output_constraints([out_sharding])
autop.optimize_placement()
data = autop.sharding_optimizer.get_json()
from autoparallel.visualizer.build_display_from_json import generate_visualization_html
with open("viz.html", "w") as f:
f.write(generate_visualization_html(data))Or from a saved optimizer:
opt = ShardingOptimizer.load("model.ap")
data = opt.get_json()save()/load(): Usestorch.save(pickle). Same-codebase, same-PyTorch-version only. Custom ops must be registered before loading (the loader importsautoparallel.cast_parametrizationautomatically).save_placements()/load_placements(): Plain JSON. Portable across runs, but the model graph and mesh must match.load_placements()validates mesh shape and dim names, and verifies that every saved node name exists in the current graph.get_json(): Requires a solved optimizer — callget_solution()oroptimize_placement()first, or load a saved optimizer that already contains a solution. The output is a Python dict, not written to disk — serialize it yourself if needed.