#include "clip.h"
#include <cstdint>
#include <functional>
#include <map>
#include <optional>
#include <string>
#include <stdexcept>
#include <thread>
#include <vector>
#include "cliprs/src/lib.rs.h"
#include "rust_interface.h"
extern "C" void clip_log_warning(const char* message) {
log_warning(rust::String(message));
}
const int verbosity = 1;
static int n_threads = 1;
static std::map<uintptr_t, int> vec_dims;
struct clip_ctx * init(rust::String model_path) {
n_threads = std::thread::hardware_concurrency();
struct clip_ctx * ctx = clip_model_load(model_path.data(), verbosity);
int vec_dim = clip_get_vision_hparams(ctx)->projection_dim;
vec_dims[reinterpret_cast<uintptr_t>(ctx)] = vec_dim;
return ctx;
}
rust::vec<float> embed_text(const struct clip_ctx * ctx,rust::String text) {
struct clip_tokens tokens;
clip_tokenize(ctx, text.c_str(), &tokens);
std::vector<float> txt_vec;
txt_vec.insert(txt_vec.end(), vec_dims[reinterpret_cast<uintptr_t>(ctx)], 0);
if (!clip_text_encode(ctx, n_threads, &tokens, txt_vec.data(), true)) {
throw std::runtime_error(std::string(__func__) + "Failed to encode text");
}
rust::vec<float> ret;
std::copy(txt_vec.begin(), txt_vec.end(), std::back_inserter(ret));
return ret;
}
rust::vec<float> embed_image(const struct clip_ctx * ctx, rust::String path) {
std::string path_str(path);
struct clip_image_u8 * img0 = clip_image_u8_make();
if (!clip_image_load_from_file(path.c_str(), img0)) {
throw std::runtime_error(std::string(__func__) + ": Failed to load image from " + path_str);
}
struct clip_image_f32 * img_res = clip_image_f32_make();
if (!clip_image_preprocess(ctx, img0, img_res)) {
throw std::runtime_error(std::string(__func__) + ": Failed to preprocess " + path_str);
}
std::vector<float> img_vec;
img_vec.insert(img_vec.end(), vec_dims[reinterpret_cast<uintptr_t>(ctx)], 0);
if (!clip_image_encode(ctx, n_threads, img_res, img_vec.data(), true)) {
throw std::runtime_error(std::string(__func__) + ": Failed to encode " + path_str);
}
rust::vec<float> ret;
std::copy(img_vec.begin(), img_vec.end(), std::back_inserter(ret));
return ret;
}
float embed_compare(const struct clip_ctx * ctx, const rust::vec<float> & p1, const rust::vec<float> & p2) {
return clip_similarity_score(p1.data(), p2.data(), vec_dims[reinterpret_cast<uintptr_t>(ctx)]);
}
void end(struct clip_ctx * ctx) { clip_free(ctx); }