bool has_vision = false;
bool has_audio = false;
+ mtmd_progress_callback progress_callback = nullptr;
+ void * progress_callback_user_data = nullptr;
+
// TODO @ngxson : we should not pass clip_ctx here, it should be clip_model
- clip_model_loader(const char * fname, bool skip_tensors = false) : fname(fname) {
+ clip_model_loader(const char * fname,
+ bool skip_tensors = false,
+ mtmd_progress_callback progress_cb = nullptr,
+ void * progress_user_data = nullptr)
+ : fname(fname),
+ progress_callback(progress_cb),
+ progress_callback_user_data(progress_user_data) {
struct ggml_context * meta = nullptr;
struct gguf_init_params params = {
if (!ctx_clip.no_alloc) {
std::vector<uint8_t> read_buf;
+ // start loading event
+ if (progress_callback){
+ progress_callback(0.0, progress_callback_user_data);
+ }
+
+ // compute total tensor data size for progress reporting
+ size_t total_data_size = 0;
+ for (auto & t : tensors_to_load) {
+ total_data_size += ggml_nbytes(t);
+ }
+
// alloc memory and offload data
ggml_backend_buffer_type_t buft = ggml_backend_get_default_buffer_type(ctx_clip.backend);
ctx_clip.buf.reset(ggml_backend_alloc_ctx_tensors_from_buft(ctx_clip.ctx_data.get(), buft));
ggml_backend_buffer_set_usage(ctx_clip.buf.get(), GGML_BACKEND_BUFFER_USAGE_WEIGHTS);
+ size_t data_loaded = 0;
for (auto & t : tensors_to_load) {
ggml_tensor * cur = ggml_get_tensor(ctx_clip.ctx_data.get(), t->name);
GGML_ASSERT(cur && "tensor not found in ctx_data");
fin.read(reinterpret_cast<char *>(read_buf.data()), num_bytes);
ggml_backend_tensor_set(cur, read_buf.data(), 0, num_bytes);
}
+ data_loaded += num_bytes;
+ if (progress_callback && total_data_size > 0) {
+ const float progress = (float)data_loaded / (float)total_data_size;
+ if (!progress_callback(progress, progress_callback_user_data)) {
+ throw std::runtime_error(string_format("%s: model loading cancelled by progress_callback\n", __func__));
+ }
+ }
}
fin.close();
clip_ctx * ctx_audio = nullptr;
try {
- clip_model_loader loader(fname);
+ clip_model_loader loader(fname,
+ /* skip_tensors */ false,
+ ctx_params.progress_callback,
+ ctx_params.progress_callback_user_data);
bool skip_audio = false;
if (loader.has_vision) {
ggml_backend_sched_eval_callback cb_eval;
void * cb_eval_user_data;
bool no_alloc;
+ mtmd_progress_callback progress_callback;
+ void * progress_callback_user_data;
};
struct clip_init_result {
/* cb_eval */ nullptr,
/* cb_eval_user_data */ nullptr,
/* batch_max_tokens */ 1024,
+ /* progress_callback */ nullptr,
+ /* progress_callback_user_data */ nullptr,
};
return params;
}
/* cb_eval */ ctx_params.cb_eval,
/* cb_eval_user_data */ ctx_params.cb_eval_user_data,
/* no_alloc */ no_alloc,
+ /* progress_callback */ ctx_params.progress_callback,
+ /* progress_callback_user_data */ ctx_params.progress_callback_user_data,
};
auto res = clip_init(mmproj_fname, ctx_clip_params);
mtmd::context_ptr ctx;
auto saved_log_callback = g_logger_state.log_callback;
auto saved_log_user_data = g_logger_state.log_callback_user_data;
+
+ ctx_params.progress_callback = nullptr;
+
try {
mtmd_log_set(stub_log_callback, nullptr); // suppress logging
+ // TODO @ngxson : fix no_alloc here
ctx.reset(new mtmd_context(mmproj_fname, nullptr, ctx_params));
mtmd_log_set(saved_log_callback, saved_log_user_data); // restore log callback
std::map<ggml_backend_dev_t, size_t> total_mem;
typedef struct mtmd_input_text mtmd_input_text;
typedef struct mtmd_batch mtmd_batch;
+typedef bool (*mtmd_progress_callback)(float progress, void * user_data);
+
struct mtmd_context_params {
bool use_gpu;
bool print_timings;
int32_t batch_max_tokens; // maximum number of output tokens in a batch
// (note: this is not a hard-limit, the first image will always be added even if it exceeds this limit)
// (default: 1024)
+
+ // Called with a progress value between 0.0 and 1.0. Pass NULL to disable.
+ // If the provided progress_callback returns true, model loading continues.
+ // If it returns false, model loading is immediately aborted.
+ mtmd_progress_callback progress_callback;
+ void * progress_callback_user_data;
};
MTMD_API const char * mtmd_default_marker(void);
bool sleeping = false;
- int64_t t_last_load_progress_ms = 0;
-
void destroy() {
spec.reset();
ctx_dft.reset();
sleeping = new_state;
}
+ struct load_progress_data {
+ server_context_impl * ctx;
+ std::string stage;
+ int64_t t_last_load_progress_ms = 0;
+ load_progress_data(server_context_impl * ctx, const std::string & stage) : ctx(ctx), stage(stage) {}
+ };
static bool load_progress_callback(float progress, void * user_data) {
- auto * ctx = static_cast<server_context_impl *>(user_data);
- GGML_ASSERT(ctx);
+ auto * d = static_cast<load_progress_data *>(user_data);
+ GGML_ASSERT(d);
// always emit the first and final sample; throttle the rest to one per 200ms
{
- auto & t_last = ctx->t_last_load_progress_ms;
+ auto & t_last = d->t_last_load_progress_ms;
const int64_t t_now = ggml_time_ms();
const bool first = t_last == 0;
const bool done = progress >= 1.0f;
}
t_last = t_now;
}
- if (ctx->callback_state) {
- ctx->callback_state(SERVER_STATE_LOADING, {
- {"stage", "text_model"},
+ if (d->ctx->callback_state) {
+ d->ctx->callback_state(SERVER_STATE_LOADING, {
+ {"stage", d->stage},
{"value", progress},
});
}
// load the model and initialize llama_context
// this may also be called to resume from sleeping state
bool load_model(common_params & params) {
+ load_progress_data load_progress_text(this, "text_model");
+ load_progress_data load_progress_mmproj(this, "mmproj_model");
+
bool is_resume = sleeping;
SRV_INF("loading model '%s'\n", params.model.path.c_str());
mparams.image_max_tokens = params_base.image_max_tokens;
mparams.batch_max_tokens = params_base.mtmd_batch_max_tokens;
mparams.media_marker = get_media_marker();
+ // progress callback
+ mparams.progress_callback = load_progress_callback;
+ mparams.progress_callback_user_data = &load_progress_mmproj;
}
// optionally get the memory usage of mmproj
// attach a progress callback
{
- t_last_load_progress_ms = 0;
params_base.load_progress_callback = load_progress_callback;
- params_base.load_progress_callback_user_data = this;
+ params_base.load_progress_callback_user_data = &load_progress_text;
}
llama_init = common_init_from_params(params_base);
}
if (has_mmproj) {
- if (callback_state) {
- callback_state(SERVER_STATE_LOADING, {{"stage", "mmproj_model"}});
- }
-
if (!is_resume) {
mtmd_helper_log_set(common_log_default_callback, nullptr);
}