Added logic to normalize tensors if needed

This commit is contained in:
cerspense
2024-02-27 11:22:12 -08:00
parent 96969b1a2f
commit 9b60341a09
+8
View File
@@ -79,6 +79,10 @@ class Modelscopet2v:
if not isinstance(video_frames, torch.Tensor):
video_frames = torch.tensor(video_frames, dtype=torch.float32)
# Normalize the tensor to have values between 0 and 1 if they are in the range 0-255
if video_frames.max() > 1.0:
video_frames = video_frames / 255.0
# Remove the unnecessary batch dimension explicitly and permute the dimensions
# The expected shape is (num_frames, height, width, channels)
video_frames = video_frames.squeeze(0).permute(0, 1, 2, 3)
@@ -140,6 +144,10 @@ class Modelscopev2v:
if not isinstance(video_frames, torch.Tensor):
video_frames = torch.tensor(video_frames, dtype=torch.float32)
# Normalize the tensor to have values between 0 and 1 if they are in the range 0-255
if video_frames.max() > 1.0:
video_frames = video_frames / 255.0
# The expected shape is (num_frames, height, width, channels)
video_frames = video_frames.squeeze(0).permute(0, 1, 2, 3)