aitune.torch.jit.patched_module
aitune.torch.jit.patched_module
Module for inspecting PyTorch models and tracking their execution.
Module Contents
Classes
Functions
Data
PRINT_HIERARCHY_NO_MODULES_HEADER
API
Bases: Exception
Exception raised when graph break is detected.
Bases: enum.Enum
Possible states of the Module class.
Patched module.
Intercepts torch.nn.Module’s forward function to record data, tune module and to do inference on a tuned one.
As opposed to the AITune declarative approach, where the user specifies which module should be tuned, JIT approach intercepts all modules on the fly and creates call hierarchy. After a few calls it can perform tuning. JIT is actually a bit harder case, since during navigation of the model hierarchy we are not aware of the real parent object (there could be python objects in between). Hence we cannot override parent object reference meaning we can’t inject module as a proxy. This is why only forward function is overridden and device attribute is patched with dynamic class creation so that only instance attribute is changed.
Get the module cache directory.
Get the fully qualified name of the module.
This is guaranteed to be unique after first forward call.
Representation of the module.
String representation of the module.
Create a cache directory for the graph.
Forward call for the first time.
This method is called when the module is first called. It initializes the module and sets the state to RECORDING.
Forward call for the recording state.
In this state we already have the hierarchy resolved for the current module. If necessary we can skip the module and its children from processing.
Route forward calls to the appropriate state-specific handler.
This method serves as a stable wrapper throughout the module’s lifecycle. This design is necessary for compatibility with models that internally cache references to module.forward wrappers. Since the wrapper object itself never changes, cached references remain valid across state transitions.
Forward call for the tuned state.
Return a unique fully qualified name for this module in the hierarchy.
Builds the path from root to this module by prepending the parent’s already-computed fq_name (e.g. grand_parent_name.parent_name.name). An index suffix is only added when there is a duplicate path (same path seen more than once).
Get the hash of the module hierarchy.
Assumption: given name (actual class name + number of parameters) + id is unique enough to be used as a hash
Handle new hooks if added by a backend.
Before tuning hooks are cleared (empty OrderedDict). If there are new hooks added by a backend, they should be added to the existing ones but before the original hooks so that application layer hooks (like HF post processing hooks) are called after the backend hooks.
Return JSON-ready graph specs recorded before deferred tuning.
Log a JIT tuning exception and persist its traceback.
Return this module’s cache directory.
Return unique parameter and buffer dtypes for this module tree.
Patch the device attribute of the module and its parents.
HF pipelines have device attribute which checks next(module.parameters()).device.
Since JIT tuning moves the original module to a cpu device, this attribute returns cpu device.
This causes misplacement issues where some modules are on cuda and other modules or tensors are moved according
to the device attribute to cpu. To fix this, the method patches dynamically device attribute but only at instance
(object) level i.e. does not touch class object. This has to be applied not only to the current module but to
its parents in case parents objects do not have it’s own parameters but include child module which has been
tuned and move to a cpu device.
Proxy the forward calls with a stable wrapper.
Replaces torch.nn.Module.forward method with a proxy forward that routes to the appropriate handler. The substitution is done with a wrapt.decorator so that the replaced function has the same docstring, signature and other attributes. This is crucial as some HF models perform self inspection for method arguments.
We re-enable hooks so that they are called before and after the proxied forward.
Restore the original forward and hooks.
We need to disable hooks, otherwise they will be called twice.
Set original forwards for all patched descendants in the module tree.
The PatchedModule child hierarchy tracks modules observed during execution. It does not include registered submodules that were never called for the recorded samples (for example eval-only Identity drop-path modules). Backends receive the full nn.Module ownership tree, and some of them deepcopy/export that tree, so restore every patched descendant before handing the module to a backend.
Check if the module should be skipped.
Name contains parameters, so only class names is relevant.
Check if the module should be tuned, dispatching based on the active JIT mode.
Check if the module should be tuned in deferred mode.
The explicit enable_tune_deferred() call marks that recording is complete, so the next normal forward tunes eligible modules after recording that call.
Check if the module should be tuned in eager mode.
Requires a minimum number of forward passes and, when dynamic shapes are expected, at least one sample with a detected dynamic axis.
Return whether this tuning path should emit per-module inspection details.
Simulate tuning in dry-run mode.
It runs tune_dry_run but it also raises exception with probability of config.dry_run_failure_probability.
This is done to simulate failures during JIT tuning to see how it would behave.
Throw if graph break is detected.
Adds to history the duration of the checking.
Tune the module with optional graph break detection and timing.
Tune the module on init.
Unpatch the module.
Removes it also from Patcher object registry.
Unpatch the module and all its children.
Update the state of the module and its forward method.
For states with custom forward handlers (INIT, RECORDING, TUNED), this method installs a proxy forward that routes to the appropriate handler. For other states, the original forward is restored.
Parameters:
The new state of the module
Persist a tuning traceback in this module’s cache directory.
Build a report snapshot of the JIT data observed for this module.
Build report snapshots for this module and observed child modules.
Give suggestions if tuning was not attempted.
Prints the PatchedModule hierarchy starting from the head module.
This method traverses the module tree starting from the head module and prints each module with indentation to show the hierarchy levels.
Prints history to the sink.
Reset PatchedModule state.
Tune the module if readiness conditions for the active JIT mode are met.
Delegates the readiness check to _should_be_tuned(), which dispatches to the mode-specific implementation (eager or deferred).
Tunes the module in blocking mode.
Parameters:
if True, only dry run the tuning
Build the per-module tune strategy.
Resolves the strategy via config.resolve_strategy() and clones it so each module gets
a fresh instance — strategies hold per-tune state (backend_results etc.) that must
not leak across modules.
Find-max-batch-size profiling is disabled because in JIT we cannot run the original module separately; the strategy must work from the recorded samples alone.
Updates history with new entry.