aitune.torch.jit.patched_module

View as Markdown

Module for inspecting PyTorch models and tracking their execution.

Module Contents

Classes

NameDescription
GraphBreakExceptionException raised when graph break is detected.
ModuleStatePossible states of the Module class.
PatchedModulePatched module.

Functions

NameDescription
_build_strategyBuild the per-module tune strategy.
_to_histUpdates history with new entry.

Data

PRINT_HIERARCHY_HEADER

PRINT_HIERARCHY_NO_MODULES_HEADER

_emoji_state

API

class aitune.torch.jit.patched_module.GraphBreakException()
Exception

Bases: Exception

Exception raised when graph break is detected.

class aitune.torch.jit.patched_module.ModuleState

Bases: enum.Enum

Possible states of the Module class.

DETACHED
= 'detached'
EAGER
= 'eager'
INIT
= 'init'
RECORDING
= 'recording'
SKIPPED
= 'skipped'
TUNED
= 'tuned'
class aitune.torch.jit.patched_module.PatchedModule(
module: torch.nn.Module
)

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.

_call_count
= 0
_children
list[PatchedModule] = []
_current_forward_hooks
= module._forward_hooks
_current_forward_pre_hooks
= module._forward_pre_hooks
_extra_state_info
str = ''
_forward_routing
_fq_name
str | None = None
_id
= PatchedModule.module_counter
_level
= -1
_name
= module.__class__.__name__
_original_forward
= module.forward
_parent
PatchedModule | None = None
attempted_tuning
bool = False
cache_dir
Path

Get the module cache directory.

deferred_tuning_enabled
bool = False
fq_name
str

Get the fully qualified name of the module.

This is guaranteed to be unique after first forward call.

fq_name_counter
Counter[str] = Counter()
heads
list[PatchedModule] = []
history
list[str] = []
module_counter
int = 0
patched_classes
Counter[str] = Counter()
stack
deque[PatchedModule] = deque()
aitune.torch.jit.patched_module.PatchedModule.__repr__()

Representation of the module.

aitune.torch.jit.patched_module.PatchedModule.__str__()

String representation of the module.

aitune.torch.jit.patched_module.PatchedModule._create_graph_cache_dir(
graph_spec_name: str
) -> pathlib.Path

Create a cache directory for the graph.

aitune.torch.jit.patched_module.PatchedModule._forward_init(
wrapped,
instance,
args,
kwargs
)

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.

aitune.torch.jit.patched_module.PatchedModule._forward_recording(
wrapped,
instance,
args,
kwargs
)

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.

aitune.torch.jit.patched_module.PatchedModule._forward_router(
wrapped,
instance,
args,
kwargs
)

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.

aitune.torch.jit.patched_module.PatchedModule._forward_tuned(
wrapped,
instance,
args,
kwargs
)

Forward call for the tuned state.

aitune.torch.jit.patched_module.PatchedModule._get_fully_qualified_name() -> str

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).

aitune.torch.jit.patched_module.PatchedModule._get_hierarchy_hash() -> str

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

aitune.torch.jit.patched_module.PatchedModule._handle_backend_added_hooks()

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.

aitune.torch.jit.patched_module.PatchedModule._inspection_graphs() -> list[dict]

Return JSON-ready graph specs recorded before deferred tuning.

aitune.torch.jit.patched_module.PatchedModule._log_tuning_exception(
exception: Exception
) -> None

Log a JIT tuning exception and persist its traceback.

aitune.torch.jit.patched_module.PatchedModule._module_cache_dir() -> pathlib.Path

Return this module’s cache directory.

aitune.torch.jit.patched_module.PatchedModule._module_dtypes() -> list[str]

Return unique parameter and buffer dtypes for this module tree.

aitune.torch.jit.patched_module.PatchedModule._patch_device_attribute(
device: torch.device
)

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.

aitune.torch.jit.patched_module.PatchedModule._proxy_forward()

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.

aitune.torch.jit.patched_module.PatchedModule._restore_original_forward()

Restore the original forward and hooks.

We need to disable hooks, otherwise they will be called twice.

aitune.torch.jit.patched_module.PatchedModule._set_original_forward_for_hierarchy()

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.

aitune.torch.jit.patched_module.PatchedModule._should_be_skipped()

Check if the module should be skipped.

Name contains parameters, so only class names is relevant.

aitune.torch.jit.patched_module.PatchedModule._should_be_tuned() -> bool

Check if the module should be tuned, dispatching based on the active JIT mode.

aitune.torch.jit.patched_module.PatchedModule._should_be_tuned_deferred() -> bool

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.

aitune.torch.jit.patched_module.PatchedModule._should_be_tuned_eager() -> bool

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.

aitune.torch.jit.patched_module.PatchedModule._should_report_inspection() -> bool

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.

aitune.torch.jit.patched_module.PatchedModule._tune_on_init()

Tune the module on init.

aitune.torch.jit.patched_module.PatchedModule._unpatch()

Unpatch the module.

Removes it also from Patcher object registry.

aitune.torch.jit.patched_module.PatchedModule._unpatch_hierarchy(
include_self = False
)

Unpatch the module and all its children.

aitune.torch.jit.patched_module.PatchedModule._update_state(
)

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:

state
ModuleState

The new state of the module

aitune.torch.jit.patched_module.PatchedModule._write_error_log(
exception: Exception
) -> pathlib.Path | None

Persist a tuning traceback in this module’s cache directory.

aitune.torch.jit.patched_module.PatchedModule.inspection_report() -> aitune.torch.tune_data.report_models.ModuleInspectionReport

Build a report snapshot of the JIT data observed for this module.

aitune.torch.jit.patched_module.PatchedModule.inspection_subtree_reports() -> list[aitune.torch.tune_data.report_models.ModuleInspectionReport]

Build report snapshots for this module and observed child modules.

aitune.torch.jit.patched_module.PatchedModule.on_python_exit()
staticmethod

Give suggestions if tuning was not attempted.

aitune.torch.jit.patched_module.PatchedModule.print_hierarchy(
sink = print
)
staticmethod

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.

aitune.torch.jit.patched_module.PatchedModule.print_history(
sink = print
)
staticmethod

Prints history to the sink.

aitune.torch.jit.patched_module.PatchedModule.reset()
staticmethod

Reset PatchedModule state.

aitune.torch.jit.patched_module.PatchedModule.try_tune()

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).

aitune.torch.jit.patched_module.PatchedModule.tune()

Tunes the module in blocking mode.

Parameters:

dry_run

if True, only dry run the tuning

aitune.torch.jit.patched_module._build_strategy() -> aitune.torch.tune_strategy.tune_strategy.TuneStrategy

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.

aitune.torch.jit.patched_module._to_hist(
entry: str
)

Updates history with new entry.

aitune.torch.jit.patched_module.PRINT_HIERARCHY_HEADER = 'JIT Tuning Hierarchy:'
aitune.torch.jit.patched_module.PRINT_HIERARCHY_NO_MODULES_HEADER = 'No modules in hierarchy'
aitune.torch.jit.patched_module._emoji_state = {ModuleState.INIT: '⏳', ModuleState.RECORDING: '🔴', ModuleState.TUNED: '🎯', Mo...