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,
)
| 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 | |
| 1279 | def 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 |