处理类型异常。
This commit is contained in:
@@ -161,10 +161,13 @@ class FoodSegmentationDataset(Dataset):
|
|||||||
# albumentations会自动将增强同时应用到image和mask
|
# albumentations会自动将增强同时应用到image和mask
|
||||||
transformed = self.transform(image=image, mask=mask)
|
transformed = self.transform(image=image, mask=mask)
|
||||||
image = transformed['image'] # Tensor (3, H, W)
|
image = transformed['image'] # Tensor (3, H, W)
|
||||||
mask = transformed['mask'] # ndarray (H, W)
|
mask = transformed['mask'] # 可能是ndarray或Tensor
|
||||||
|
|
||||||
# 3. 将mask转换为Tensor
|
# 3. 将mask转换为Tensor(检查类型)
|
||||||
mask = torch.from_numpy(mask).long()
|
if isinstance(mask, np.ndarray):
|
||||||
|
mask = torch.from_numpy(mask).long()
|
||||||
|
else:
|
||||||
|
mask = mask.long() # 已经是Tensor,直接转换类型
|
||||||
|
|
||||||
# 4. 检查mask的有效性
|
# 4. 检查mask的有效性
|
||||||
# 确保mask的值在[0, num_classes-1]范围内
|
# 确保mask的值在[0, num_classes-1]范围内
|
||||||
|
|||||||
Reference in New Issue
Block a user