(carry_i: NestedTensor, scan_i: NestedTensor)
| 1333 | representing the `fn` outputs and output collection of the ith scan iteration, |
| 1334 | respesctively. |
| 1335 | |
| 1336 | Raises: |
| 1337 | ValueError: If `current_context()` is None, or if invalid remat_kwargs are passed. |
| 1338 | """ |
| 1339 | |
| 1340 | ctx = current_context() |
| 1341 | if ctx is None: |
| 1342 | raise ValueError("Expected current_context() to not be None.") |
| 1343 | |
| 1344 | def scan_fn(carry_i: NestedTensor, scan_i: NestedTensor): |
| 1345 | output_collection_i = new_output_collection() |
| 1346 | x_i = scan_i.pop("xs") |
| 1347 | with child_context( |
| 1348 | "iter", |
| 1349 | module=ctx.module, |
| 1350 | output_collection=output_collection_i, |
| 1351 | **scan_i, |
| 1352 | ): |
| 1353 | carry_i, y_i = fn(carry_i, x_i) |
| 1354 | |
| 1355 | # Filter output collection. |
| 1356 | if drop_output is not None: |
| 1357 | pruned_collection_i = new_output_collection()._asdict() |
| 1358 | pruned_collection_i.update( |
| 1359 | prune_tree( |
| 1360 | output_collection_i._asdict(), |
nothing calls this directly
no test coverage detected