Recursively run collate logic and provide detailed loggings for debugging purposes. It reports results at the 'critical' level, is therefore suitable in the context of exception handling. Args: batch: batch input to collate level: current level of recursion for logging
(batch, level: int = 1, logger_name: str = "dev_collate")
| 353 | |
| 354 | |
| 355 | def dev_collate(batch, level: int = 1, logger_name: str = "dev_collate"): |
| 356 | """ |
| 357 | Recursively run collate logic and provide detailed loggings for debugging purposes. |
| 358 | It reports results at the 'critical' level, is therefore suitable in the context of exception handling. |
| 359 | |
| 360 | Args: |
| 361 | batch: batch input to collate |
| 362 | level: current level of recursion for logging purposes |
| 363 | logger_name: name of logger to use for logging |
| 364 | |
| 365 | See also: https://pytorch.org/docs/stable/data.html#working-with-collate-fn |
| 366 | """ |
| 367 | elem = batch[0] |
| 368 | elem_type = type(elem) |
| 369 | l_str = ">" * level |
| 370 | batch_str = f"{batch[:10]}{' ... ' if len(batch) > 10 else ''}" |
| 371 | if isinstance(elem, torch.Tensor): |
| 372 | try: |
| 373 | logging.getLogger(logger_name).critical(f"{l_str} collate/stack a list of tensors") |
| 374 | return torch.stack(batch, 0) |
| 375 | except TypeError as e: |
| 376 | logging.getLogger(logger_name).critical( |
| 377 | f"{l_str} E: {e}, type {[type(elem).__name__ for elem in batch]} in collate({batch_str})" |
| 378 | ) |
| 379 | return |
| 380 | except RuntimeError as e: |
| 381 | logging.getLogger(logger_name).critical( |
| 382 | f"{l_str} E: {e}, shape {[elem.shape for elem in batch]} in collate({batch_str})" |
| 383 | ) |
| 384 | return |
| 385 | elif elem_type.__module__ == "numpy" and elem_type.__name__ != "str_" and elem_type.__name__ != "string_": |
| 386 | if elem_type.__name__ in ["ndarray", "memmap"]: |
| 387 | logging.getLogger(logger_name).critical(f"{l_str} collate/stack a list of numpy arrays") |
| 388 | return dev_collate([torch.as_tensor(b) for b in batch], level=level, logger_name=logger_name) |
| 389 | elif elem.shape == (): # scalars |
| 390 | return batch |
| 391 | elif isinstance(elem, (float, int, str, bytes)): |
| 392 | return batch |
| 393 | elif isinstance(elem, abc.Mapping): |
| 394 | out = {} |
| 395 | for key in elem: |
| 396 | logging.getLogger(logger_name).critical(f'{l_str} collate dict key "{key}" out of {len(elem)} keys') |
| 397 | out[key] = dev_collate([d[key] for d in batch], level=level + 1, logger_name=logger_name) |
| 398 | return out |
| 399 | elif isinstance(elem, abc.Sequence): |
| 400 | it = iter(batch) |
| 401 | els = list(it) |
| 402 | try: |
| 403 | sizes = [len(elem) for elem in els] # may not have `len` |
| 404 | except TypeError: |
| 405 | types = [type(elem).__name__ for elem in els] |
| 406 | logging.getLogger(logger_name).critical(f"{l_str} E: type {types} in collate({batch_str})") |
| 407 | return |
| 408 | logging.getLogger(logger_name).critical(f"{l_str} collate list of sizes: {sizes}.") |
| 409 | if any(s != sizes[0] for s in sizes): |
| 410 | logging.getLogger(logger_name).critical( |
| 411 | f"{l_str} collate list inconsistent sizes, got size: {sizes}, in collate({batch_str})" |
| 412 | ) |
searching dependent graphs…