Skip to content

PyTorch

Torch frontend for ZKDV.

ZKDV

ZKDV(path: Any, config: ZKDVConfig | None = None, *, max_in_flight: int = 2, replay_snapshot_interval: int = 12)

Bases: Driver

Generate a verification tape around ordinary Torch training steps.

Source code in zkdv/driver/torch.py
def __init__(
    self,
    path: Any,
    config: ZKDVConfig | None = None,
    *,
    max_in_flight: int = 2,
    replay_snapshot_interval: int = 12,
) -> None:
    super().__init__(
        path,
        config,
        backend="torch",
        annotation=torch.profiler.record_function,
        max_in_flight=max_in_flight,
        replay_snapshot_interval=replay_snapshot_interval,
    )
    self._attested_function = None
    self._lowered_program = None
    self._prepared_model: PreparedModel | None = None
    self._prepared_optimizer: PreparedOptimizer | None = None
    self._optimizer_registry = OptimizerRegistry()
    self._prepared_compiled = None
    self._prepared_requires_deltas = False

backward

backward(loss: Tensor, gradient: Any = None) -> None

Lower backward in attestation and use eager autograd in training.

Source code in zkdv/driver/torch.py
def backward(self, loss: torch.Tensor, gradient: Any = None) -> None:
    """Lower backward in attestation and use eager autograd in training."""

    capture = current_capture()
    if capture is not None:
        capture.backward(loss, gradient)
    else:
        torch.autograd.backward(loss, gradient)

forward

forward(model: Module, *args: Any, **kwargs: Any) -> Any

Lower a prepared model call during attestation capture.

Source code in zkdv/driver/torch.py
def forward(self, model: torch.nn.Module, *args: Any, **kwargs: Any) -> Any:
    """Lower a prepared model call during attestation capture."""

    capture = current_capture()
    if capture is None:
        return model(*args, **kwargs)
    prepared = self._require_prepared_model()
    if model is not prepared.public:
        raise RuntimeError(
            "proof.forward() received an unprepared model; "
            "pass it through proof.prepare() first"
        )
    return prepared.forward(capture.model_state, args, kwargs)

prepare

prepare(*values: Any) -> Any

Prepare recognized stateful objects and leave all others unchanged.

Source code in zkdv/driver/torch.py
def prepare(self, *values: Any) -> Any:
    """Prepare recognized stateful objects and leave all others unchanged."""

    prepared = tuple(self._prepare(value) for value in values)
    return prepared[0] if len(prepared) == 1 else prepared

register_functional_optimizer

register_functional_optimizer(optimizer_type: type[Optimizer], contract: OptimizerContract) -> None

Register a proof-local custom functional optimizer contract.

Source code in zkdv/driver/torch.py
def register_functional_optimizer(
    self,
    optimizer_type: type[torch.optim.Optimizer],
    contract: OptimizerContract,
) -> None:
    """Register a proof-local custom functional optimizer contract."""

    self._optimizer_registry.register(contract, optimizer_type)

transaction

transaction(*, batch: Any, overlap: Any = None) -> EagerTransaction

Begin an eager transaction before the model forward mutates state.

Source code in zkdv/driver/torch.py
def transaction(self, *, batch: Any, overlap: Any = None) -> EagerTransaction:
    """Begin an eager transaction before the model forward mutates state."""

    if overlap is None:
        from zkdv import overlap as overlap_policy

        overlap = overlap_policy.UNCHECKED
    return EagerTransaction(self, batch, overlap)

artifact

Hermetic Torch AOT artifact compilation and loading.

roundtrip

roundtrip(function: Any, arguments: tuple[Any, ...]) -> tuple[Any, bytearray]

Compile, serialize, and reload one exact AOT Inductor program.

Source code in zkdv/torch/artifact.py
def roundtrip(function: Any, arguments: tuple[Any, ...]) -> tuple[Any, bytearray]:
    """Compile, serialize, and reload one exact AOT Inductor program."""

    with torch._dynamo.config.patch(trace_autograd_ops=True):
        compiled = torch.compile(function, fullgraph=True).aot_compile((arguments, {}))
    with tempfile.TemporaryDirectory() as directory:
        artifact = pathlib.Path(directory, "compiled.pt")
        compiled.save_compiled_function(str(artifact))
        material = bytearray(artifact.read_bytes())
    return load(material), material

capture

Proof-aware execution context for lowering an attested Torch transition.

AttestationCapture

AttestationCapture(model, model_state, optimizer, optimizer_state)

Collect gradients and the optimizer transition during attested execution.

Source code in zkdv/torch/capture.py
def __init__(self, model, model_state, optimizer, optimizer_state) -> None:
    self.model = model
    self._buffers_before = model_state[1]
    self.model_state = (
        model_state[0],
        tuple(buffer.clone() for buffer in model_state[1]),
    )
    self.optimizer = optimizer
    self.optimizer_state = optimizer_state
    self.gradients = None
    self.transition = None

checkpoint

Bounded pinned-host snapshots for Torch checked replay.

eager

One eager Torch training transaction submitted through the ZKDV pipeline.

EagerTransaction

EagerTransaction(proof, batch, overlap=UNCHECKED)

Keep a proof transaction open from eager forward through optimizer step.

Source code in zkdv/torch/eager.py
def __init__(self, proof, batch, overlap=overlap_policy.UNCHECKED) -> None:
    if _CURRENT.get() is not None:
        raise RuntimeError("nested ZKDV Torch transactions are unsupported")
    self.proof = proof
    self.batch = batch
    self.overlap = overlap
    self.buffers_before = tuple(
        buffer.clone() for buffer in proof._require_prepared_model().buffers
    )
    self.pending = proof._open_prepared_transaction(self)
    self._token: Token | None = _CURRENT.set(self)
    self._complete = False

generated

Import target used by packaged ZKDV Torch AOT guards.

lowering

Lower proof-aware Python execution to a standalone Torch graph function.

lower

lower(function: Any, arguments: tuple[Any, ...]) -> Any

Capture once and return a state-free callable for trusted Rust sealing.

Source code in zkdv/torch/lowering.py
def lower(function: Any, arguments: tuple[Any, ...]) -> Any:
    """Capture once and return a state-free callable for trusted Rust sealing."""

    torch._dynamo.reset()
    try:
        with torch._dynamo.config.patch(trace_autograd_ops=True):
            graph = torch._dynamo.export(
                function,
                aten_graph=True,
                same_signature=False,
            )(*arguments).graph_module
    except torch._dynamo.exc.Unsupported as error:
        cause = error
        while cause is not None:
            message = str(cause)
            if "pass it through proof.prepare() first" in message:
                raise RuntimeError(message) from error
            cause = cause.__cause__
        raise
    captured = tuple(graph.named_parameters()) + tuple(graph.named_buffers())
    if captured:
        names = ", ".join(name for name, _ in captured)
        raise RuntimeError(
            f"attested Torch code captured unprepared tensor state ({names}); "
            "pass the owning object through proof.prepare()"
        )
    random_operations = tuple(
        node.target._schema.name
        for node in graph.graph.nodes
        if node.op == "call_function"
        and getattr(node.target, "_schema", None) is not None
        and any(
            marker in node.target._schema.name
            for marker in (
                "bernoulli",
                "dropout",
                "exponential",
                "multinomial",
                "normal",
                "rand",
                "uniform",
            )
        )
    )
    if random_operations:
        raise RuntimeError(
            f"attested Torch code contains random operations "
            f"{random_operations}; generate randomness outside the transition "
            "and include it in the prepared transaction batch"
        )
    source = graph.code.replace("def forward(self, ", "def lowered(")
    generated = __import__("zkdv.torch.generated", fromlist=("lowered",))
    namespace = generated.__dict__
    namespace.update(graph.forward.__globals__)
    namespace["__name__"] = generated.__name__
    exec(compile(source, "<zkdv-torch-attestation>", "exec"), namespace)
    flat_arguments = tuple(value.detach() for value in pytree.tree_leaves(arguments))
    with torch.no_grad():
        _, material = roundtrip(namespace["lowered"], flat_arguments)
    return material

model

Prepared model state used by Torch functional replay.

PreparedModel dataclass

PreparedModel(public: Module, target: Module, parameter_names: tuple[str, ...], buffer_names: tuple[str, ...], distributed: bool, fsdp: bool, parameter_metadata: tuple[Any, ...], state_parameters: tuple[Tensor, ...], state_buffers: tuple[Tensor, ...])

Bind one eager module to a stable parameter and buffer layout.

optimizers

Functional optimizer contracts used by the Torch frontend.

OptimizerBridge

OptimizerBridge(optimizer: Optimizer, contract: OptimizerContract)

Own the canonical state layout for one prepared optimizer.

Source code in zkdv/torch/optimizers/bridge.py
def __init__(
    self,
    optimizer: torch.optim.Optimizer,
    contract: OptimizerContract,
) -> None:
    contract.validate(optimizer)
    self.optimizer = optimizer
    self.contract = contract
    self.parameters = tuple(
        parameter
        for group in optimizer.param_groups
        for parameter in group["params"]
    )
    self._group_sizes = tuple(
        len(group["params"]) for group in optimizer.param_groups
    )
    group_positions = []
    start = 0
    for size in self._group_sizes:
        group_positions.append(tuple(range(start, start + size)))
        start += size
    self._group_positions = tuple(group_positions)
    self._configuration = self._current_configuration()
    self._initialize()
    pack_optimizer_state(optimizer)

bind

bind(parameters: tuple[Tensor, ...]) -> None

Bind optimizer groups to a model's canonical parameter order.

Source code in zkdv/torch/optimizers/bridge.py
def bind(self, parameters: tuple[torch.Tensor, ...]) -> None:
    """Bind optimizer groups to a model's canonical parameter order."""

    positions = {id(parameter): index for index, parameter in enumerate(parameters)}
    if len(positions) != len(parameters) or set(positions) != {
        id(parameter) for parameter in self.parameters
    }:
        raise ValueError(
            "prepared optimizer parameters must exactly cover the prepared model"
        )
    self._group_positions = tuple(
        tuple(positions[id(parameter)] for parameter in group["params"])
        for group in self.optimizer.param_groups
    )

pack_state

pack_state() -> tuple[tuple[dict[str, Any], ...], ...]

Return the live optimizer state without copying its tensors.

Source code in zkdv/torch/optimizers/bridge.py
def pack_state(self) -> tuple[tuple[dict[str, Any], ...], ...]:
    """Return the live optimizer state without copying its tensors."""

    return tuple(
        tuple(
            {
                name: value.to_local() if isinstance(value, DTensor) else value
                for name, value in self.optimizer.state[parameter].items()
            }
            for parameter in group["params"]
        )
        for group in self.optimizer.param_groups
    )

transition

transition(parameters: tuple[Tensor, ...], state: tuple[tuple[dict[str, Any], ...], ...], gradients: tuple[Tensor | None, ...]) -> tuple[tuple[Tensor, ...], tuple[tuple[dict[str, Any], ...], ...]]

Compute deltas and next optimizer state without changing the inputs.

Source code in zkdv/torch/optimizers/bridge.py
def transition(
    self,
    parameters: tuple[torch.Tensor, ...],
    state: tuple[tuple[dict[str, Any], ...], ...],
    gradients: tuple[torch.Tensor | None, ...],
) -> tuple[tuple[torch.Tensor, ...], tuple[tuple[dict[str, Any], ...], ...]]:
    """Compute deltas and next optimizer state without changing the inputs."""

    deltas = [None] * len(parameters)
    states_after = []
    for group, positions, group_state in zip(
        self.optimizer.param_groups,
        self._group_positions,
        state,
        strict=True,
    ):
        current = tuple(parameters[position] for position in positions)
        current_gradients = tuple(gradients[position] for position in positions)
        updated = [parameter.clone() for parameter in current]
        updated_state = [clone_state(value) for value in group_state]
        active = [
            position
            for position, gradient in enumerate(current_gradients)
            if gradient is not None
        ]
        with torch.no_grad():
            if active:
                self.contract.update(
                    [updated[position] for position in active],
                    [current_gradients[position] for position in active],
                    [updated_state[position] for position in active],
                    group,
                )
            for position, before, after in zip(
                positions,
                current,
                updated,
                strict=True,
            ):
                deltas[position] = after - before
        states_after.append(tuple(updated_state))
    if any(delta is None for delta in deltas):
        raise RuntimeError("optimizer bridge left a model parameter unbound")
    return tuple(deltas), tuple(states_after)

validate_configuration

validate_configuration() -> None

Reject optimizer semantics that changed after attestation.

Source code in zkdv/torch/optimizers/bridge.py
def validate_configuration(self) -> None:
    """Reject optimizer semantics that changed after attestation."""

    if self._current_configuration() != self._configuration:
        raise RuntimeError(
            "prepared optimizer options changed after attestation; create a "
            "new proof for the new optimizer schedule"
        )

OptimizerContract

Bases: ABC

Describe one optimizer's tensor state and functional group update.

initialize abstractmethod

initialize(parameter: Tensor, group: dict[str, Any]) -> dict

Create the state Torch would lazily initialize for one parameter.

Source code in zkdv/torch/optimizers/base.py
@abstractmethod
def initialize(self, parameter: torch.Tensor, group: dict[str, Any]) -> dict:
    """Create the state Torch would lazily initialize for one parameter."""

update abstractmethod

update(parameters: list[Tensor], gradients: list[Tensor], states: list[dict[str, Any]], group: dict[str, Any]) -> None

Apply one functional parameter-group update in place.

Source code in zkdv/torch/optimizers/base.py
@abstractmethod
def update(
    self,
    parameters: list[torch.Tensor],
    gradients: list[torch.Tensor],
    states: list[dict[str, Any]],
    group: dict[str, Any],
) -> None:
    """Apply one functional parameter-group update in place."""

validate

validate(optimizer: Optimizer) -> None

Reject modes whose semantics cannot be represented by the bridge.

Source code in zkdv/torch/optimizers/base.py
def validate(self, optimizer: torch.optim.Optimizer) -> None:
    """Reject modes whose semantics cannot be represented by the bridge."""

    for group in optimizer.param_groups:
        if group.get("differentiable", False):
            raise ValueError(
                "ZKDV does not support differentiable optimizer steps; "
                "register a custom functional optimizer contract"
            )
        tensor_options = tuple(
            name
            for name, value in group.items()
            if name != "params" and isinstance(value, torch.Tensor)
        )
        if tensor_options:
            raise ValueError(
                f"ZKDV requires static optimizer options, but "
                f"{tensor_options} are tensors; register a contract that "
                "includes dynamic options in optimizer state"
            )

OptimizerRegistry

OptimizerRegistry()

Resolve prepared optimizers without optimizer branching elsewhere.

Source code in zkdv/torch/optimizers/registry.py
def __init__(self) -> None:
    self._contracts: dict[type[torch.optim.Optimizer], OptimizerContract] = {}
    for contract in (
        ASGDContract(),
        AdadeltaContract(),
        AdafactorContract(),
        AdagradContract(),
        AdamContract(),
        AdamWContract(),
        AdamaxContract(),
        LBFGSContract(),
        MuonContract(),
        NAdamContract(),
        RAdamContract(),
        RMSpropContract(),
        RpropContract(),
        SGDContract(),
        SparseAdamContract(),
    ):
        self.register(contract)

adadelta

Functional bridge contract for :class:torch.optim.Adadelta.

adafactor

Functional bridge contract for :class:torch.optim.Adafactor.

adagrad

Functional bridge contract for :class:torch.optim.Adagrad.

adam

Functional bridge contract for :class:torch.optim.Adam.

adamax

Functional bridge contract for :class:torch.optim.Adamax.

adamw

Functional bridge contract for :class:torch.optim.AdamW.

asgd

Functional bridge contract for :class:torch.optim.ASGD.

base

Contract between an eager Torch optimizer and its functional update.

OptimizerContract

Bases: ABC

Describe one optimizer's tensor state and functional group update.

initialize abstractmethod

initialize(parameter: Tensor, group: dict[str, Any]) -> dict

Create the state Torch would lazily initialize for one parameter.

Source code in zkdv/torch/optimizers/base.py
@abstractmethod
def initialize(self, parameter: torch.Tensor, group: dict[str, Any]) -> dict:
    """Create the state Torch would lazily initialize for one parameter."""

update abstractmethod

update(parameters: list[Tensor], gradients: list[Tensor], states: list[dict[str, Any]], group: dict[str, Any]) -> None

Apply one functional parameter-group update in place.

Source code in zkdv/torch/optimizers/base.py
@abstractmethod
def update(
    self,
    parameters: list[torch.Tensor],
    gradients: list[torch.Tensor],
    states: list[dict[str, Any]],
    group: dict[str, Any],
) -> None:
    """Apply one functional parameter-group update in place."""

validate

validate(optimizer: Optimizer) -> None

Reject modes whose semantics cannot be represented by the bridge.

Source code in zkdv/torch/optimizers/base.py
def validate(self, optimizer: torch.optim.Optimizer) -> None:
    """Reject modes whose semantics cannot be represented by the bridge."""

    for group in optimizer.param_groups:
        if group.get("differentiable", False):
            raise ValueError(
                "ZKDV does not support differentiable optimizer steps; "
                "register a custom functional optimizer contract"
            )
        tensor_options = tuple(
            name
            for name, value in group.items()
            if name != "params" and isinstance(value, torch.Tensor)
        )
        if tensor_options:
            raise ValueError(
                f"ZKDV requires static optimizer options, but "
                f"{tensor_options} are tensors; register a contract that "
                "includes dynamic options in optimizer state"
            )

bridge

State-preserving bridge from eager optimizers to functional contracts.

OptimizerBridge

OptimizerBridge(optimizer: Optimizer, contract: OptimizerContract)

Own the canonical state layout for one prepared optimizer.

Source code in zkdv/torch/optimizers/bridge.py
def __init__(
    self,
    optimizer: torch.optim.Optimizer,
    contract: OptimizerContract,
) -> None:
    contract.validate(optimizer)
    self.optimizer = optimizer
    self.contract = contract
    self.parameters = tuple(
        parameter
        for group in optimizer.param_groups
        for parameter in group["params"]
    )
    self._group_sizes = tuple(
        len(group["params"]) for group in optimizer.param_groups
    )
    group_positions = []
    start = 0
    for size in self._group_sizes:
        group_positions.append(tuple(range(start, start + size)))
        start += size
    self._group_positions = tuple(group_positions)
    self._configuration = self._current_configuration()
    self._initialize()
    pack_optimizer_state(optimizer)

bind

bind(parameters: tuple[Tensor, ...]) -> None

Bind optimizer groups to a model's canonical parameter order.

Source code in zkdv/torch/optimizers/bridge.py
def bind(self, parameters: tuple[torch.Tensor, ...]) -> None:
    """Bind optimizer groups to a model's canonical parameter order."""

    positions = {id(parameter): index for index, parameter in enumerate(parameters)}
    if len(positions) != len(parameters) or set(positions) != {
        id(parameter) for parameter in self.parameters
    }:
        raise ValueError(
            "prepared optimizer parameters must exactly cover the prepared model"
        )
    self._group_positions = tuple(
        tuple(positions[id(parameter)] for parameter in group["params"])
        for group in self.optimizer.param_groups
    )

pack_state

pack_state() -> tuple[tuple[dict[str, Any], ...], ...]

Return the live optimizer state without copying its tensors.

Source code in zkdv/torch/optimizers/bridge.py
def pack_state(self) -> tuple[tuple[dict[str, Any], ...], ...]:
    """Return the live optimizer state without copying its tensors."""

    return tuple(
        tuple(
            {
                name: value.to_local() if isinstance(value, DTensor) else value
                for name, value in self.optimizer.state[parameter].items()
            }
            for parameter in group["params"]
        )
        for group in self.optimizer.param_groups
    )

transition

transition(parameters: tuple[Tensor, ...], state: tuple[tuple[dict[str, Any], ...], ...], gradients: tuple[Tensor | None, ...]) -> tuple[tuple[Tensor, ...], tuple[tuple[dict[str, Any], ...], ...]]

Compute deltas and next optimizer state without changing the inputs.

Source code in zkdv/torch/optimizers/bridge.py
def transition(
    self,
    parameters: tuple[torch.Tensor, ...],
    state: tuple[tuple[dict[str, Any], ...], ...],
    gradients: tuple[torch.Tensor | None, ...],
) -> tuple[tuple[torch.Tensor, ...], tuple[tuple[dict[str, Any], ...], ...]]:
    """Compute deltas and next optimizer state without changing the inputs."""

    deltas = [None] * len(parameters)
    states_after = []
    for group, positions, group_state in zip(
        self.optimizer.param_groups,
        self._group_positions,
        state,
        strict=True,
    ):
        current = tuple(parameters[position] for position in positions)
        current_gradients = tuple(gradients[position] for position in positions)
        updated = [parameter.clone() for parameter in current]
        updated_state = [clone_state(value) for value in group_state]
        active = [
            position
            for position, gradient in enumerate(current_gradients)
            if gradient is not None
        ]
        with torch.no_grad():
            if active:
                self.contract.update(
                    [updated[position] for position in active],
                    [current_gradients[position] for position in active],
                    [updated_state[position] for position in active],
                    group,
                )
            for position, before, after in zip(
                positions,
                current,
                updated,
                strict=True,
            ):
                deltas[position] = after - before
        states_after.append(tuple(updated_state))
    if any(delta is None for delta in deltas):
        raise RuntimeError("optimizer bridge left a model parameter unbound")
    return tuple(deltas), tuple(states_after)

validate_configuration

validate_configuration() -> None

Reject optimizer semantics that changed after attestation.

Source code in zkdv/torch/optimizers/bridge.py
def validate_configuration(self) -> None:
    """Reject optimizer semantics that changed after attestation."""

    if self._current_configuration() != self._configuration:
        raise RuntimeError(
            "prepared optimizer options changed after attestation; create a "
            "new proof for the new optimizer schedule"
        )

buffers

Contiguous storage for native optimizer state tensors.

pack_optimizer_state

pack_optimizer_state(optimizer: Optimizer) -> None

Rebind compatible state tensors to shared contiguous backing buffers.

Source code in zkdv/torch/optimizers/buffers.py
def pack_optimizer_state(optimizer: torch.optim.Optimizer) -> None:
    """Rebind compatible state tensors to shared contiguous backing buffers."""

    groups = {}
    for parameter, state in optimizer.state.items():
        for name, value in state.items():
            if (
                not isinstance(value, torch.Tensor)
                or isinstance(value, DTensor)
                or value.layout is not torch.strided
            ):
                continue
            key = (
                name,
                tuple(value.shape),
                value.dtype,
                value.device,
                value.requires_grad,
            )
            groups.setdefault(key, []).append((parameter, value))

    with torch.no_grad():
        for (name, shape, dtype, device, requires_grad), entries in groups.items():
            if len(entries) < 2:
                continue
            buffer = torch.empty(
                (len(entries), *shape),
                dtype=dtype,
                device=device,
                requires_grad=requires_grad,
            )
            for position, (parameter, value) in enumerate(entries):
                buffer[position].copy_(value)
                optimizer.state[parameter][name] = buffer[position]

lbfgs

Explicit rejection contract for closure-driven LBFGS updates.

muon

Functional bridge contract for :class:torch.optim.Muon.

nadam

Functional bridge contract for :class:torch.optim.NAdam.

radam

Functional bridge contract for :class:torch.optim.RAdam.

registry

Exact-type registry for Torch functional optimizer contracts.

OptimizerRegistry

OptimizerRegistry()

Resolve prepared optimizers without optimizer branching elsewhere.

Source code in zkdv/torch/optimizers/registry.py
def __init__(self) -> None:
    self._contracts: dict[type[torch.optim.Optimizer], OptimizerContract] = {}
    for contract in (
        ASGDContract(),
        AdadeltaContract(),
        AdafactorContract(),
        AdagradContract(),
        AdamContract(),
        AdamWContract(),
        AdamaxContract(),
        LBFGSContract(),
        MuonContract(),
        NAdamContract(),
        RAdamContract(),
        RMSpropContract(),
        RpropContract(),
        SGDContract(),
        SparseAdamContract(),
    ):
        self.register(contract)

rmsprop

Functional bridge contract for :class:torch.optim.RMSprop.

rprop

Functional bridge contract for :class:torch.optim.Rprop.

sgd

Functional bridge contract for :class:torch.optim.SGD.

sparse_adam

Functional bridge contract for :class:torch.optim.SparseAdam.

utils

Flat tensor utilities shared by functional optimizer contracts.

pipeline

Canonical Torch transaction pipeline above the native protocol boundary.

Pipeline

Pipeline(core: Any, program: Any, queue: Any, params: Any, opt_state: Any, batch: Any, *, snapshot_interval: int = 12, deferred_mutation_barrier: bool = False)

Bases: Pipeline

Compile and submit the Torch fused update protocol.

Source code in zkdv/torch/pipeline.py
def __init__(
    self,
    core: Any,
    program: Any,
    queue: Any,
    params: Any,
    opt_state: Any,
    batch: Any,
    *,
    snapshot_interval: int = 12,
    deferred_mutation_barrier: bool = False,
) -> None:
    placement = pytree.tree_leaves(params)[0].device
    (
        self._batch_to_host,
        self._start,
        self._check,
        self._register,
        self._abort,
        self._compile,
        self._prepare,
        self._resolve,
        checked_program,
        program_material,
    ) = core.pipeline(program, params, opt_state, str(placement), batch)
    self._update_checks_possible = core.pipeline_update_checks_possible()
    self._device = placement
    core.program_artifact(program_material, params, opt_state)
    checkpoint = ReplayCheckpoint(
        params,
        opt_state,
        batch,
        checked_program,
        snapshot_interval=snapshot_interval,
        snapshot_pool_size=queue.capacity,
        deferred_mutation_barrier=deferred_mutation_barrier,
    )
    super().__init__(core, queue, checkpoint)

requires_update_deltas property

requires_update_deltas: bool

Whether a sampled replay can require production update deltas.

prepared

Eager-facing objects returned by :meth:zkdv.torch.ZKDV.prepare.

PreparedOptimizer

PreparedOptimizer(proof, optimizer, bridge)

Bases: Optimizer

Preserve the Optimizer API while routing steps through ZKDV.

Source code in zkdv/torch/prepared.py
def __init__(self, proof, optimizer, bridge) -> None:
    self.proof = proof
    self.optimizer = optimizer
    self.bridge = bridge