Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 12 additions & 5 deletions src/models/qwen_vl_vision.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -249,6 +249,8 @@ std::vector<float> QwenVisionPipeline::Run(const float* pixel_data, const std::v
// Matches HuggingFace transformers implementation:
// https://github.com/huggingface/transformers/blob/main/src/transformers/models/qwen2_5_vl/modeling_qwen2_5_vl.py#L367
std::vector<int64_t> QwenVisionPipeline::CalculateWindowIndex(int64_t grid_t, int64_t grid_h, int64_t grid_w) {
ValidateWindowIndexParams(grid_t, grid_h, grid_w, spatial_merge_size_, patch_size_, window_size_);

// Calculate LLM grid dimensions after spatial merging
int64_t llm_grid_h = grid_h / spatial_merge_size_;
int64_t llm_grid_w = grid_w / spatial_merge_size_;
Expand All @@ -260,21 +262,26 @@ std::vector<int64_t> QwenVisionPipeline::CalculateWindowIndex(int64_t grid_t, in
int64_t pad_h = (vit_merger_window_size - (llm_grid_h % vit_merger_window_size)) % vit_merger_window_size;
int64_t pad_w = (vit_merger_window_size - (llm_grid_w % vit_merger_window_size)) % vit_merger_window_size;

int64_t num_windows_h = (llm_grid_h + pad_h) / vit_merger_window_size;
int64_t num_windows_w = (llm_grid_w + pad_w) / vit_merger_window_size;
int64_t padded_h = llm_grid_h + pad_h;
int64_t padded_w = llm_grid_w + pad_w;

int64_t alloc_size = grid_t * padded_h * padded_w;

int64_t num_windows_h = padded_h / vit_merger_window_size;
int64_t num_windows_w = padded_w / vit_merger_window_size;

std::vector<int64_t> window_index;
window_index.reserve(grid_t * llm_grid_h * llm_grid_w);

// Create initial index grid
std::vector<int64_t> index(grid_t * (llm_grid_h + pad_h) * (llm_grid_w + pad_w), -100);
std::vector<int64_t> index(alloc_size, -100);

// Fill non-padded positions with sequential indices
for (int64_t t = 0; t < grid_t; ++t) {
for (int64_t h = 0; h < llm_grid_h; ++h) {
for (int64_t w = 0; w < llm_grid_w; ++w) {
int64_t idx = t * llm_grid_h * llm_grid_w + h * llm_grid_w + w;
int64_t padded_idx = t * (llm_grid_h + pad_h) * (llm_grid_w + pad_w) + h * (llm_grid_w + pad_w) + w;
int64_t padded_idx = t * padded_h * padded_w + h * padded_w + w;
index[padded_idx] = idx;
}
}
Expand All @@ -290,7 +297,7 @@ std::vector<int64_t> QwenVisionPipeline::CalculateWindowIndex(int64_t grid_t, in
for (int64_t pw = 0; pw < vit_merger_window_size; ++pw) {
int64_t h = wh * vit_merger_window_size + ph;
int64_t w = ww * vit_merger_window_size + pw;
int64_t padded_idx = t * (llm_grid_h + pad_h) * (llm_grid_w + pad_w) + h * (llm_grid_w + pad_w) + w;
int64_t padded_idx = t * padded_h * padded_w + h * padded_w + w;

// Only add non-padded indices
if (index[padded_idx] != -100) {
Expand Down
34 changes: 34 additions & 0 deletions src/models/qwen_vl_vision.h
Original file line number Diff line number Diff line change
Expand Up @@ -7,11 +7,45 @@
#include <vector>
#include <memory>
#include <cstdint>
#include <stdexcept>

#include "onnxruntime_api.h"

namespace Generators {

// Validates grid dimensions and config parameters for the Qwen vision window indexing.
// Throws std::runtime_error on invalid inputs.
// Defined inline in the header so that unit tests can call it directly without requiring
// DLL export, since this is a free function not part of the public C API.
inline void ValidateWindowIndexParams(int64_t grid_t, int64_t grid_h, int64_t grid_w,
int64_t spatial_merge_size, int64_t patch_size, int64_t window_size) {
if (spatial_merge_size <= 0)
throw std::runtime_error("CalculateWindowIndex: spatial_merge_size must be positive");
if (patch_size <= 0)
throw std::runtime_error("CalculateWindowIndex: patch_size must be positive");
if (grid_t <= 0 || grid_h <= 0 || grid_w <= 0)
throw std::runtime_error("CalculateWindowIndex: grid dimensions must be positive");
if (grid_h % spatial_merge_size != 0 || grid_w % spatial_merge_size != 0)
throw std::runtime_error("CalculateWindowIndex: grid_h and grid_w must be divisible by spatial_merge_size");

int64_t vit_merger_window_size = window_size / spatial_merge_size / patch_size;
if (vit_merger_window_size <= 0)
throw std::runtime_error("CalculateWindowIndex: vit_merger_window_size must be positive (check window_size, spatial_merge_size, patch_size config)");

constexpr int64_t kMaxElements = static_cast<int64_t>(1) << 30;
int64_t llm_grid_h = grid_h / spatial_merge_size;
int64_t llm_grid_w = grid_w / spatial_merge_size;
if (llm_grid_h > kMaxElements || llm_grid_w > kMaxElements || grid_t > kMaxElements)
throw std::runtime_error("CalculateWindowIndex: grid dimensions are too large");

int64_t pad_h = (vit_merger_window_size - (llm_grid_h % vit_merger_window_size)) % vit_merger_window_size;
int64_t pad_w = (vit_merger_window_size - (llm_grid_w % vit_merger_window_size)) % vit_merger_window_size;
int64_t padded_h = llm_grid_h + pad_h;
int64_t padded_w = llm_grid_w + pad_w;
if (padded_h > 0 && padded_w > 0 && grid_t > kMaxElements / padded_h / padded_w)
throw std::runtime_error("CalculateWindowIndex: total grid size exceeds maximum allowed");
}

// Internal vision pipeline (no external DLL interface required after Python binding removal).
struct QwenVisionPipeline {
QwenVisionPipeline(OrtEnv& env,
Expand Down
45 changes: 44 additions & 1 deletion test/model_tests.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,9 @@
#include <ort_genai.h>
#include <gtest/gtest.h>

#include "models/model.h"
#include "models/qwen_vl_vision.h"

#include "test_utils.h"

// External global variable from main.cpp for custom model path
Expand Down Expand Up @@ -477,4 +480,44 @@ Print all primes between 1 and n

std::cout << tokenizer->Decode(result) << "\r\n";
}
#endif
#endif

// --- Validation tests (no model files required) ---

TEST(ValidationTests, WindowIndexAcceptsValidParams) {
EXPECT_NO_THROW(Generators::ValidateWindowIndexParams(1, 28, 28, 2, 14, 112));
EXPECT_NO_THROW(Generators::ValidateWindowIndexParams(1, 2, 2, 2, 14, 56));
}

TEST(ValidationTests, WindowIndexRejectsZeroDivisors) {
EXPECT_THROW(Generators::ValidateWindowIndexParams(1, 28, 28, 0, 14, 112), std::runtime_error);
EXPECT_THROW(Generators::ValidateWindowIndexParams(1, 28, 28, -1, 14, 112), std::runtime_error);
EXPECT_THROW(Generators::ValidateWindowIndexParams(1, 28, 28, 2, 0, 112), std::runtime_error);
EXPECT_THROW(Generators::ValidateWindowIndexParams(1, 28, 28, 2, -14, 112), std::runtime_error);
}

TEST(ValidationTests, WindowIndexRejectsInvalidGridDims) {
EXPECT_THROW(Generators::ValidateWindowIndexParams(0, 28, 28, 2, 14, 112), std::runtime_error);
EXPECT_THROW(Generators::ValidateWindowIndexParams(1, -28, 28, 2, 14, 112), std::runtime_error);
EXPECT_THROW(Generators::ValidateWindowIndexParams(1, 28, 0, 2, 14, 112), std::runtime_error);
}

TEST(ValidationTests, WindowIndexRejectsNonDivisibleGrid) {
EXPECT_THROW(Generators::ValidateWindowIndexParams(1, 27, 28, 2, 14, 112), std::runtime_error);
EXPECT_THROW(Generators::ValidateWindowIndexParams(1, 28, 27, 2, 14, 112), std::runtime_error);
}

TEST(ValidationTests, WindowIndexRejectsZeroMergerWindowSize) {
// window_size=1 / spatial_merge_size=2 / patch_size=14 = 0
EXPECT_THROW(Generators::ValidateWindowIndexParams(1, 28, 28, 2, 14, 1), std::runtime_error);
}

TEST(ValidationTests, WindowIndexRejectsExcessiveDimensions) {
int64_t huge = (static_cast<int64_t>(1) << 31) + 2; // After /2, llm_grid_h > kMaxElements (1<<30)
EXPECT_THROW(Generators::ValidateWindowIndexParams(1, huge, 2, 2, 14, 112), std::runtime_error);
}

TEST(ValidationTests, WindowIndexRejectsTotalSizeOverflow) {
// Each dim individually <= kMaxElements, but grid_t * padded_h * padded_w > kMaxElements
EXPECT_THROW(Generators::ValidateWindowIndexParams(1000, 2000000, 2000000, 2, 14, 112), std::runtime_error);
}
Loading