MCPcopy Create free account
hub / github.com/Project-MONAI/MONAI / dev_collate

Function dev_collate

monai/data/utils.py:355–416  ·  view source on GitHub ↗

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")

Source from the content-addressed store, hash-verified

353
354
355def 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 )

Callers 2

test_dev_collateMethod · 0.90
list_data_collateFunction · 0.85

Calls 1

as_tensorMethod · 0.80

Tested by 1

test_dev_collateMethod · 0.72

Used in the wild real call sites across dependent graphs

searching dependent graphs…