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

Function scan_fn

axlearn/common/module.py:1335–1357  ·  view source on GitHub ↗
(carry_i: NestedTensor, scan_i: NestedTensor)

Source from the content-addressed store, hash-verified

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(),

Callers

nothing calls this directly

Calls 6

prune_treeFunction · 0.90
new_output_collectionFunction · 0.85
child_contextFunction · 0.85
OutputCollectionClass · 0.85
fnFunction · 0.70
updateMethod · 0.45

Tested by

no test coverage detected