Fix ImageListSampler to unwrap batches and handle list inputs correctly

This commit is contained in:
Marco
2026-04-10 23:38:33 -03:00
parent 9c7bdf8636
commit b4f448eabb
+16 -8
View File
@@ -23,22 +23,30 @@ class ImageListSampler:
if not images:
return ([],)
total = len(images)
# Flatten list/batches into individual frames
import torch
all_frames = []
for item in images:
if isinstance(item, torch.Tensor):
if len(item.shape) == 4: # [B, H, W, C]
for i in range(item.shape[0]):
all_frames.append(item[i:i+1])
else:
all_frames.append(item.unsqueeze(0) if len(item.shape) == 3 else item)
elif isinstance(item, list):
all_frames.extend(item)
total = len(all_frames)
if cnt <= 0:
return ([],)
if cnt >= total:
return (images,)
return (all_frames,)
# Calculate equally spaced indices
# Example: total=9, count=3
# step = 9 / 3 = 3.0
# i=0 -> 0
# i=1 -> 3
# i=2 -> 6
step = total / cnt
indices = [int(i * step) for i in range(cnt)]
sampled = [images[i] for i in indices]
sampled = [all_frames[i] for i in indices]
return (sampled,)