MCPcopy Create free account
hub / github.com/apple/axlearn / scan_in_context

Function scan_in_context

axlearn/common/module.py:1270–1377  ·  view source on GitHub ↗

A thin wrapper around `jax.lax.scan` which is compatible with `OutputCollection`. In particular, summaries and outputs added by `add_summary` and `add_module_output` respectively are accumulated in `current_context().output_collection`, subject to any output filtering. Specifically, sum

(
    fn,
    *,
    carry: NestedTensor,
    xs: NestedTensor,
    drop_output: Optional[Callable[[str], bool]] = None,
    child_name_prefix: str = "iter",
    unroll: Union[int, bool] = 1,
    remat_kwargs: Optional[dict[str, Any]] = None,
    merge_summaries: bool = False,
)

Source from the content-addressed store, hash-verified

1268
1269 fn = _Functional(
1270 context=context, method_fn=method_fn, require_parent=True, copy_args_tree=copy_args_tree
1271 )
1272 method_outputs, output_collection = fn(*args, **kwargs)
1273
1274 for output_collection_type in drop_output_collections:
1275 getattr(output_collection, output_collection_type).clear()
1276 return method_outputs, output_collection
1277
1278
1279def scan_in_context(
1280 fn,
1281 *,
1282 carry: NestedTensor,
1283 xs: NestedTensor,
1284 drop_output: Optional[Callable[[str], bool]] = None,
1285 child_name_prefix: str = "iter",
1286 unroll: Union[int, bool] = 1,
1287 remat_kwargs: Optional[dict[str, Any]] = None,
1288 merge_summaries: bool = False,
1289) -> tuple[NestedTensor, NestedTensor]:
1290 """A thin wrapper around `jax.lax.scan` which is compatible with `OutputCollection`.
1291
1292 In particular, summaries and outputs added by `add_summary` and `add_module_output` respectively
1293 are accumulated in `current_context().output_collection`, subject to any output filtering.
1294 Specifically, summaries from iteration `i` will be placed in
1295 `summaries[f"{child_name_prefix}{i}"]`; module outputs will be stacked and placed in
1296 `module_outputs[child_name_prefix]`.
1297
1298 Args:
1299 fn: A function with args (carry, x) returning a dict(carry=..., y=...).
1300 carry: The initial value of the loop carry, to be accumulated across scan.
1301 xs: A dict with at least "x" as a key, where each leaf is a tensor of shape
1302 [num_scan_iters, ...]. At scan iteration i:
1303 - xs["x"][i, ...] represents the inputs to `fn`.
1304 - xs[key][i, ...] is provided as a kwarg to the ith invocation context.
1305 drop_output: A callable that takes a path and outputs a decision of whether to drop the
1306 output at the given path, where True means we drop. By default, the callable is None,
1307 meaning nothing is dropped.
1308 child_name_prefix: The child name prefix used for children to be added to
1309 `target_output_collection`.
1310 unroll: If a positive integer is provided, it determines how many unrolled loop iterations
1311 to run within a single rolled iteration of the loop. If a boolean is provided, it will
1312 determine if the loop is competely unrolled (i.e. unroll=True) or left completely rolled
1313 (i.e. unroll=False).
1314 remat_kwargs: Optional dict passed to `jax.checkpoint` to enable rematerialization.
1315 Common options include:
1316 - `prevent_cse`: (bool) Whether to disable common subexpression elimination.
1317 If left unspecified, defaults to False following recommendations from the JAX
1318 documentation.
1319 Raises a ValueError if `prevent_cse` is set to True.
1320 - `policy`: A checkpoint policy from `jax.checkpoint_policies`.
1321 If provided, the scan body will be wrapped as:
1322 `scan_fn = jax.checkpoint(scan_fn, **remat_kwargs)`
1323 Otherwise, `jax.checkpoint` is not used.
1324 See https://docs.jax.dev/en/latest/_autosummary/jax.checkpoint.html.
1325 merge_summaries: If True, accumulate summaries across scan iterations via
1326 Summary.accumulate() instead of creating per-iteration children.
1327

Callers 3

_chunked_metricsMethod · 0.90
_runMethod · 0.90
_invokeMethod · 0.90

Calls 3

current_contextFunction · 0.85
scanMethod · 0.45

Tested by 1

_invokeMethod · 0.72