File size: 4,857 Bytes
d2816aa | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 | #include "testing.h"
#include "mtmd-image.h"
#include "mtmd-internal.h"
#include <iostream>
#include <stdexcept>
#include <string>
#include <tuple>
#include <utility>
#include <vector>
// this test file contains:
// 1. test cases for mtmd helpers
// 2. test cases for internal mtmd components
// internal headers can be included here
struct test_registry {
using fn_t = void (*)(testing &);
struct entry {
std::string name;
fn_t fn;
};
static std::vector<entry> & all() {
static std::vector<entry> entries;
return entries;
}
test_registry(const char * name, fn_t fn) {
all().push_back({ name, fn });
}
};
#define MAKE_TEST(name) \
static void name(testing & t); \
static const test_registry test_registry_ ## name(#name, &name); \
static void name(testing & t)
//
// mtmd_image
//
MAKE_TEST(test_image_preprocessor_lfm2) {
clip_hparams hparams;
hparams.patch_size = 16;
hparams.n_merge = 2;
hparams.set_limit_image_tokens(64, 256);
// { image size, expected tiling }
const std::vector<std::pair<clip_image_size, bool>> cases = {
{ { 704, 704 }, false },
// 720 / (patch_size * n_merge) is exactly 22.5, so this only matches HF
// if round_by_factor rounds half to even (22) instead of away from zero (23)
{ { 720, 720 }, false },
{ { 736, 736 }, true },
{ { 1024, 977 }, true },
{ { 1056, 384 }, false },
};
for (const auto & [size, expected] : cases) {
const bool actual = mtmd_image_preprocessor_lfm2::should_tile(hparams, size);
t.assert_equal(
"tiling for " + std::to_string(size.width) + "x" + std::to_string(size.height),
std::string(expected ? "tiled" : "single"),
std::string(actual ? "tiled" : "single"));
}
}
//
// mtmd temporal merge
//
MAKE_TEST(test_temporal_merge_grouping) {
std::vector<mtmd::bitmap_ptr> pool; // keeps the bitmaps alive until the end of the test
// spec chars:
// v = video frame, w = video frame of another size, a = audio, i = plain image, t = text
auto make_parts = [&pool](const std::string & spec) {
std::vector<mtmd_internal_part> parts;
for (char c : spec) {
if (c == 't') {
parts.push_back({ "hello", nullptr });
continue;
}
mtmd_bitmap * bm = nullptr;
switch (c) {
case 'v': bm = mtmd_bitmap_init(100, 100, nullptr); break;
case 'w': bm = mtmd_bitmap_init(200, 200, nullptr); break;
case 'a': bm = mtmd_bitmap_init_from_audio(100, nullptr); break;
case 'i': bm = mtmd_bitmap_init(100, 100, nullptr); break;
default: throw std::runtime_error(std::string("unknown spec char: ") + c);
}
mtmd_bitmap_set_mergeable(bm, c != 'i');
pool.emplace_back(bm);
parts.push_back({ "", bm });
}
return parts;
};
// { parts, n_merge, expected size of each group }
const std::vector<std::tuple<std::string, int, std::string>> cases = {
{ "vv", 2, "2" },
{ "vvv", 2, "21" },
{ "vvvv", 2, "22" },
{ "vvi", 2, "21" },
{ "tvvt", 2, "2" },
{ "vtv", 2, "11" }, // text in between breaks the merge
{ "vw", 2, "11" }, // different sizes cannot be merged
{ "aa", 2, "11" }, // audio is never merged
{ "ii", 2, "11" }, // two unrelated images must stay separated
{ "iv", 2, "11" },
{ "vi", 2, "11" },
{ "vv", 1, "11" }, // model without temporal merge
};
for (const auto & [spec, n_merge, expected] : cases) {
auto parts = make_parts(spec);
auto groups = mtmd_group_mergeable_bitmaps(parts, n_merge);
std::string actual;
for (const auto & group : groups) {
actual += std::to_string(group.size());
}
const std::string name = "\"" + spec + "\" with n_merge=" + std::to_string(n_merge);
t.assert_equal("groups for " + name, expected, actual);
size_t n_bitmap_parts = 0;
for (const auto & p : parts) {
n_bitmap_parts += p.bitmap != nullptr ? 1 : 0;
}
t.assert_equal("remaining bitmap parts for " + name, groups.size(), n_bitmap_parts);
}
}
//
// main
//
int main(int argc, char ** argv) {
testing t(std::cout);
t.verbose = true;
// usage: test-mtmd-impl [filter_regex]
for (int i = 1; i < argc; i++) {
t.set_filter(argv[i]);
}
for (const auto & e : test_registry::all()) {
t.test(e.name, e.fn);
}
return t.summary();
}
|