2025-12-12 21:14:48 +01:00
#pragma once
#include "ggml.h"
#include "ggml-cpp.h"
#include "clip.h"
#include "clip-impl.h"
#include "clip-model.h"
#include <vector>
#include <functional>
2025-12-16 11:25:26 +01:00
#define DEFAULT_INTERPOLATION_MODE (GGML_SCALE_MODE_BILINEAR | GGML_SCALE_FLAG_ANTIALIAS)
2026-05-15 19:32:47 +02:00
struct build_vit_opts {
ggml_tensor * attn_mask = nullptr ;
};
2025-12-12 21:14:48 +01:00
struct clip_graph {
const clip_model & model ;
const clip_hparams & hparams ;
projector_type proj_type ;
// we only support single image per batch
const clip_image_f32 & img ;
const int patch_size ;
const int n_patches_x ;
const int n_patches_y ;
const int n_patches ;
const int n_embd ;
const int n_head ;
const int d_head ;
const int n_layer ;
const int n_mmproj_embd ;
const float eps ;
2026-04-02 17:10:32 +02:00
float kq_scale ; // TODO: maybe move this to hparams
2025-12-12 21:14:48 +01:00
const clip_flash_attn_type flash_attn_type ;
ggml_context_ptr ctx0_ptr ;
ggml_context * ctx0 ;
ggml_cgraph * gf ;
clip_graph ( clip_ctx * ctx , const clip_image_f32 & img );
virtual ~ clip_graph () = default ;
virtual ggml_cgraph * build () = 0 ;
2026-03-19 13:11:39 +01:00
// wrapper around ggml_mul_mat, allow hooking (e.g. LoRA, clamping) depending on the model
// tensor w should be the weight matrix, and tensor x should be the input
virtual ggml_tensor * build_mm ( ggml_tensor * w , ggml_tensor * x ) const ;
// TODO: build_mm(w, b, x) to support bias
2025-12-12 21:14:48 +01:00
//
// utility functions
//
void cb ( ggml_tensor * cur0 , const char * name , int il ) const ;
// siglip2 naflex
2025-12-16 11:25:26 +01:00
ggml_tensor * resize_position_embeddings ( uint32_t interpolation_mode = DEFAULT_INTERPOLATION_MODE );
2025-12-12 21:14:48 +01:00
// build vision transformer (ViT) cgraph
// this function should cover most of the models
// if your model has specific features, you should probably duplicate this function
ggml_tensor * build_vit (
ggml_tensor * inp ,
int64_t n_pos ,
norm_type norm_t ,
ffn_op_type ffn_t ,
ggml_tensor * learned_pos_embd ,
2026-05-15 19:32:47 +02:00
std :: function < ggml_tensor * ( ggml_tensor * , const clip_layer & ) > add_pos ,
const build_vit_opts & opts = {});
2025-12-12 21:14:48 +01:00
// build the input after conv2d (inp_raw --> patches)
// returns tensor with shape [n_embd, n_patches]
ggml_tensor * build_inp ();
ggml_tensor * build_inp_raw ( int channels = 3 );
ggml_tensor * build_norm (
ggml_tensor * cur ,
ggml_tensor * mw ,
ggml_tensor * mb ,
norm_type type ,
float norm_eps ,
int il ) const ;
ggml_tensor * build_ffn (
ggml_tensor * cur ,
ggml_tensor * up ,
ggml_tensor * up_b ,
ggml_tensor * gate ,
ggml_tensor * gate_b ,
ggml_tensor * down ,
ggml_tensor * down_b ,
ffn_op_type type_op ,
int il ) const ;
ggml_tensor * build_attn (
ggml_tensor * wo ,
ggml_tensor * wo_b ,
ggml_tensor * q_cur ,
ggml_tensor * k_cur ,
ggml_tensor * v_cur ,
ggml_tensor * kq_mask ,
float kq_scale ,
2026-05-12 02:11:14 -07:00
int il ,
ggml_tensor * sinks = nullptr ) const ;
2025-12-12 21:14:48 +01:00
// implementation of the 2D RoPE without adding a new op in ggml
// this is not efficient (use double the memory), but works on all backends
// TODO: there was a more efficient which relies on ggml_view and ggml_rope_ext_inplace, but the rope inplace does not work well with non-contiguous tensors ; we should fix that and revert back to the original implementation in https://github.com/ggml-org/llama.cpp/pull/13065
ggml_tensor * build_rope_2d (
ggml_context * ctx0 ,
ggml_tensor * cur ,
ggml_tensor * pos_a , // first half
ggml_tensor * pos_b , // second half
const float freq_base ,
const bool interleave_freq
);
// aka pixel_shuffle / pixel_unshuffle / patch_merger (Kimi-VL)
// support dynamic resolution
ggml_tensor * build_patch_merge_permute ( ggml_tensor * cur , int scale_factor );
2025-12-15 10:18:46 +08:00
// Generic function to stack frames for audio processing
// Abstracts out the StackAudioFrames logic used by ultravox
ggml_tensor * build_stack ( ggml_tensor * cur , int32_t stack_factor , int32_t n_embed );
2025-12-12 21:14:48 +01:00
};