From 9b60341a095eefd66b4bb24c5f0a2c41771cf66a Mon Sep 17 00:00:00 2001 From: cerspense Date: Tue, 27 Feb 2024 11:22:12 -0800 Subject: [PATCH] Added logic to normalize tensors if needed --- cspnodes.py | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/cspnodes.py b/cspnodes.py index 736ff12..0be09e4 100644 --- a/cspnodes.py +++ b/cspnodes.py @@ -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)