Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 5 additions & 3 deletions agml/data/hf_loader.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@
from typing import List, Union

try:
from datasets import load_dataset, DatasetDict, Image, Sequence, ClassLabel
from datasets import load_dataset, DatasetDict, Image, Sequence, ClassLabel, Array2D
except ImportError:
raise ImportError(
"The `datasets` library is required to use the HuggingFaceDataLoader. "
Expand Down Expand Up @@ -80,8 +80,10 @@ def _cast_single(self, ds):
if "image" in features and not isinstance(features["image"], Image):
ds = ds.cast_column("image", Image())

# A "mask" column is always a pixel map — cast unconditionally.
if "mask" in features and not isinstance(features["mask"], Image):
# A "mask" column is a pixel map for image datasets, but a per-point
# label array (Array2D) for point cloud datasets. Only cast to
# Image() when it isn't already a numeric array type.
if "mask" in features and not isinstance(features["mask"], (Image, Array2D)):
ds = ds.cast_column("mask", Image())

return ds
Expand Down
Loading