mtmd: add chunk save/load function (#26645)
* mtmd: add chunk save/load function * nits * add tests * rn _MAX --> _COUNT
This commit is contained in:
@@ -1,4 +1,6 @@
|
||||
#include <stdio.h>
|
||||
#include <stdlib.h>
|
||||
#include <string.h>
|
||||
#include <assert.h>
|
||||
|
||||
#include "mtmd.h"
|
||||
@@ -62,6 +64,72 @@ int main(void) {
|
||||
}
|
||||
}
|
||||
|
||||
// test chunk save/load round-trip
|
||||
for (size_t i = 0; i < n_chunks; i++) {
|
||||
const mtmd_input_chunk * chunk = mtmd_input_chunks_get(chunks, i);
|
||||
assert(chunk != NULL);
|
||||
enum mtmd_input_chunk_type type = mtmd_input_chunk_get_type(chunk);
|
||||
|
||||
// query the required buffer size (out_buf == NULL)
|
||||
size_t expected_len = 0;
|
||||
int32_t rc = mtmd_input_chunk_save(chunk, NULL, 0, &expected_len);
|
||||
printf(" Chunk %zu: save query rc = %d, expected_len = %zu\n", i, rc, expected_len);
|
||||
assert(rc == 0);
|
||||
assert(expected_len > 0);
|
||||
|
||||
// saving into a too-small buffer must fail, not crash
|
||||
char tiny_buf[1];
|
||||
rc = mtmd_input_chunk_save(chunk, tiny_buf, sizeof(tiny_buf), NULL);
|
||||
printf(" Chunk %zu: save into too-small buffer rc = %d (expect non-zero)\n", i, rc);
|
||||
assert(rc != 0);
|
||||
|
||||
// save into a properly-sized buffer
|
||||
char * buf = (char *) malloc(expected_len);
|
||||
assert(buf != NULL);
|
||||
rc = mtmd_input_chunk_save(chunk, buf, expected_len, NULL);
|
||||
assert(rc == 0);
|
||||
|
||||
// loading from a truncated buffer must fail gracefully, not crash
|
||||
if (expected_len > 1) {
|
||||
mtmd_input_chunk * bad = mtmd_input_chunk_load(buf, expected_len - 1);
|
||||
printf(" Chunk %zu: load from truncated buffer = %p (expect NULL)\n", i, (void *) bad);
|
||||
assert(bad == NULL);
|
||||
}
|
||||
|
||||
// load it back
|
||||
mtmd_input_chunk * loaded = mtmd_input_chunk_load(buf, expected_len);
|
||||
assert(loaded != NULL);
|
||||
|
||||
// metadata must match the original chunk
|
||||
assert(mtmd_input_chunk_get_type(loaded) == type);
|
||||
assert(mtmd_input_chunk_get_n_tokens(loaded) == mtmd_input_chunk_get_n_tokens(chunk));
|
||||
assert(mtmd_input_chunk_get_n_pos(loaded) == mtmd_input_chunk_get_n_pos(chunk));
|
||||
|
||||
if (type == MTMD_INPUT_CHUNK_TYPE_TEXT) {
|
||||
size_t n_tok_orig, n_tok_loaded;
|
||||
const llama_token * tok_orig = mtmd_input_chunk_get_tokens_text(chunk, &n_tok_orig);
|
||||
const llama_token * tok_loaded = mtmd_input_chunk_get_tokens_text(loaded, &n_tok_loaded);
|
||||
printf(" Chunk %zu: loaded %zu text tokens (orig %zu), first token %d (orig %d)\n",
|
||||
i, n_tok_loaded, n_tok_orig,
|
||||
n_tok_loaded > 0 ? tok_loaded[0] : -1,
|
||||
n_tok_orig > 0 ? tok_orig[0] : -1);
|
||||
assert(n_tok_orig == n_tok_loaded);
|
||||
for (size_t j = 0; j < n_tok_orig; j++) {
|
||||
assert(tok_orig[j] == tok_loaded[j]);
|
||||
}
|
||||
} else if (type == MTMD_INPUT_CHUNK_TYPE_IMAGE || type == MTMD_INPUT_CHUNK_TYPE_AUDIO) {
|
||||
const char * id_orig = mtmd_input_chunk_get_id(chunk);
|
||||
const char * id_loaded = mtmd_input_chunk_get_id(loaded);
|
||||
printf(" Chunk %zu: loaded id '%s' (orig '%s')\n", i, id_loaded, id_orig);
|
||||
assert(id_orig != NULL && id_loaded != NULL);
|
||||
assert(strcmp(id_orig, id_loaded) == 0);
|
||||
}
|
||||
|
||||
mtmd_input_chunk_free(loaded);
|
||||
free(buf);
|
||||
}
|
||||
printf("Chunk save/load round-trip OK\n");
|
||||
|
||||
// Free the chunks
|
||||
mtmd_input_chunks_free(chunks);
|
||||
|
||||
|
||||
Reference in New Issue
Block a user