MCPcopy Create free account
hub / github.com/open-mmlab/mmengine / parameter_count_table

Function parameter_count_table

mmengine/analysis/complexity_analysis.py:295–357  ·  view source on GitHub ↗

Format the parameter count of the model (and its submodules or parameters) Adopted from https://github.com/facebookresearch/fvcore/blob/main/fvcore/nn/parameter_count.py Args: model (nn.Module): the model to count parameters. max_depth (int): maximum depth to recurs

(model: nn.Module, max_depth: int = 3)

Source from the content-addressed store, hash-verified

293
294
295def parameter_count_table(model: nn.Module, max_depth: int = 3) -> str:
296 """Format the parameter count of the model (and its submodules or
297 parameters)
298
299 Adopted from
300 https://github.com/facebookresearch/fvcore/blob/main/fvcore/nn/parameter_count.py
301
302 Args:
303 model (nn.Module): the model to count parameters.
304 max_depth (int): maximum depth to recursively print submodules or
305 parameters
306
307 Returns:
308 str: the table to be printed
309 """
310 count: typing.DefaultDict[str, int] = parameter_count(model)
311 # pyre-fixme[24]: Generic type `tuple` expects at least 1 type parameter.
312 param_shape: typing.Dict[str, typing.Tuple] = {
313 k: tuple(v.shape)
314 for k, v in model.named_parameters()
315 }
316
317 # pyre-fixme[24]: Generic type `tuple` expects at least 1 type parameter.
318 rows: typing.List[typing.Tuple] = []
319
320 def format_size(x: int) -> str:
321 if x > 1e8:
322 return f'{x / 1e9:.1f}G'
323 if x > 1e5:
324 return f'{x / 1e6:.1f}M'
325 if x > 1e2:
326 return f'{x / 1e3:.1f}K'
327 return str(x)
328
329 def fill(lvl: int, prefix: str) -> None:
330 if lvl >= max_depth:
331 return
332 for name, v in count.items():
333 if name.count('.') == lvl and name.startswith(prefix):
334 indent = ' ' * (lvl + 1)
335 if name in param_shape:
336 rows.append(
337 (indent + name, indent + str(param_shape[name])))
338 else:
339 rows.append((indent + name, indent + format_size(v)))
340 fill(lvl + 1, name + '.')
341
342 rows.append(('model', format_size(count.pop(''))))
343 fill(0, '')
344
345 table = Table(
346 title=f'parameter count of {model.__class__.__name__}', box=box.ASCII2)
347 table.add_column('name')
348 table.add_column('#elements or shape')
349
350 for row in rows:
351 table.add_row(*row)
352

Callers 1

Calls 5

parameter_countFunction · 0.85
format_sizeFunction · 0.85
fillFunction · 0.70
popMethod · 0.45
getMethod · 0.45

Tested by 1

Used in the wild real call sites across dependent graphs

searching dependent graphs…