From 881bbbf6c9d948e1bfacbabc8940a9201875dff1 Mon Sep 17 00:00:00 2001 From: sko00o Date: Thu, 31 Jul 2025 15:01:48 +0800 Subject: [PATCH] Added validation for `max_size` parameter in `get_3d_rotary_pos_embed` function when `grid_type` is set to 'slice'. --- embeddings.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/embeddings.py b/embeddings.py index 618f478..fb993d3 100644 --- a/embeddings.py +++ b/embeddings.py @@ -174,6 +174,8 @@ def get_3d_rotary_pos_embed( grid_t = np.arange(temporal_size, dtype=np.float32) grid_t = np.linspace(0, temporal_size, temporal_size, endpoint=False, dtype=np.float32) elif grid_type == "slice": + if max_size is None: + raise ValueError("`max_size` must be provided when `grid_type` is 'slice'") max_h, max_w = max_size grid_size_h, grid_size_w = grid_size grid_h = np.arange(max_h, dtype=np.float32)