cellflow.networks.ResNetBlock.tabulate¶
- ResNetBlock.tabulate(rngs, *args, depth=None, show_repeated=False, mutable=DenyList(deny='intermediates'), console_kwargs=None, table_kwargs=mappingproxy({}), column_kwargs=mappingproxy({}), compute_flops=False, compute_vjp_flops=False, **kwargs)¶
Creates a summary of the Module represented as a table.
This method has the same signature and internally calls
Module.init, but instead of returning the variables, it returns the string summarizing the Module in a table.tabulateusesjax.eval_shapeto run the forward computation without consuming any FLOPs or allocating memory.Additional arguments can be passed into the
console_kwargsargument, for example,{'width': 120}. For a full list ofconsole_kwargsarguments, see: https://rich.readthedocs.io/en/stable/reference/console.html#rich.console.ConsoleExample:
>>> import flax.linen as nn >>> import jax, jax.numpy as jnp >>> class Foo(nn.Module): ... @nn.compact ... def __call__(self, x): ... h = nn.Dense(4)(x) ... return nn.Dense(2)(h) >>> x = jnp.ones((16, 9)) >>> # print(Foo().tabulate( >>> # jax.random.key(0), x, compute_flops=True, compute_vjp_flops=True))
This gives the following output:
Foo Summary ┏━━━━━━━━━┳━━━━━━━━┳━━━━━━━━━━━━━━━┳━━━━━━━━━━━━━━━┳━━━━━━━┳━━━━━━━━━━━┳━━━━━━━━━━━━━━━━━┓ ┃ path ┃ module ┃ inputs ┃ outputs ┃ flops ┃ vjp_flops ┃ params ┃ ┡━━━━━━━━━╇━━━━━━━━╇━━━━━━━━━━━━━━━╇━━━━━━━━━━━━━━━╇━━━━━━━╇━━━━━━━━━━━╇━━━━━━━━━━━━━━━━━┩ │ │ Foo │ float32[16,9] │ float32[16,2] │ 1504 │ 4460 │ │ ├─────────┼────────┼───────────────┼───────────────┼───────┼───────────┼─────────────────┤ │ Dense_0 │ Dense │ float32[16,9] │ float32[16,4] │ 1216 │ 3620 │ bias: │ │ │ │ │ │ │ │ float32[4] │ │ │ │ │ │ │ │ kernel: │ │ │ │ │ │ │ │ float32[9,4] │ │ │ │ │ │ │ │ │ │ │ │ │ │ │ │ 40 (160 B) │ ├─────────┼────────┼───────────────┼───────────────┼───────┼───────────┼─────────────────┤ │ Dense_1 │ Dense │ float32[16,4] │ float32[16,2] │ 288 │ 840 │ bias: │ │ │ │ │ │ │ │ float32[2] │ │ │ │ │ │ │ │ kernel: │ │ │ │ │ │ │ │ float32[4,2] │ │ │ │ │ │ │ │ │ │ │ │ │ │ │ │ 10 (40 B) │ ├─────────┼────────┼───────────────┼───────────────┼───────┼───────────┼─────────────────┤ │ │ │ │ │ │ Total │ 50 (200 B) │ └─────────┴────────┴───────────────┴───────────────┴───────┴───────────┴─────────────────┘ Total Parameters: 50 (200 B)Note: rows order in the table does not represent execution order, instead it aligns with the order of keys in
variableswhich are sorted alphabetically.Note:
vjp_flopsreturns0if the module is not differentiable.- Return type:
- Parameters:
- Args:
rngs: The rngs for the variable collections as passed to
Module.init. *args: The arguments to the forward computation. depth: controls how many submodule deep the summary can go. By default,its
Nonewhich means no limit. If a submodule is not shown because of the depth limit, its parameter count and bytes will be added to the row of its first shown ancestor such that the sum of all rows always adds up to the total number of parameters of the Module.- show_repeated: If
True, repeated calls to the same module will be shown in the table, otherwise only the first call will be shown. Default is
False.- mutable: Can be bool, str, or list. Specifies which collections should be
treated as mutable:
bool: all/no collections are mutable.str: The name of a single mutable collection.list: A list of names of mutable collections. By default, all collections except ‘intermediates’ are mutable.- console_kwargs: An optional dictionary with additional keyword arguments
that are passed to
rich.console.Consolewhen rendering the table. Default arguments are{'force_terminal': True, 'force_jupyter': False}.- table_kwargs: An optional dictionary with additional keyword arguments
that are passed to
rich.table.Tableconstructor.- column_kwargs: An optional dictionary with additional keyword arguments
that are passed to
rich.table.Table.add_columnwhen adding columns to the table.- compute_flops: whether to include a
flopscolumn in the table listing the estimated FLOPs cost of each module forward pass. Does incur actual on-device computation / compilation / memory allocation, but still introduces overhead for large modules (e.g. extra 20 seconds for a Stable Diffusion’s UNet, whereas otherwise tabulation would finish in 5 seconds).
- compute_vjp_flops: whether to include a
vjp_flopscolumn in the table listing the estimated FLOPs cost of each module backward pass. Introduces a compute overhead of about 2-3X of
compute_flops.
**kwargs: keyword arguments to pass to the forward computation.
- show_repeated: If
- Returns:
A string summarizing the Module.