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)
| 293 | |
| 294 | |
| 295 | def 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 |
searching dependent graphs…