@@ -727,6 +727,14 @@ class llama_batch(ctypes.Structure):
727727LLAMA_MODEL_META_KEY_SAMPLING_MIROSTAT_ETA = 11
728728
729729
730+ # enum llama_process_type {
731+ # LLAMA_PROCESS_TYPE_ENCODE,
732+ # LLAMA_PROCESS_TYPE_DECODE,
733+ # };
734+ LLAMA_PROCESS_TYPE_ENCODE = 0
735+ LLAMA_PROCESS_TYPE_DECODE = 1
736+
737+
730738# struct llama_model_kv_override {
731739# enum llama_model_kv_override_type tag;
732740
@@ -1344,7 +1352,6 @@ def llama_ftype_name(ftype: int, /) -> Optional[bytes]:
13441352
13451353
13461354# // Initialize the llama + ggml backend
1347- # // If numa is true, use NUMA optimizations
13481355# // Call once at the start of the program
13491356# LLAMA_API void llama_backend_init(void);
13501357@ctypes_function (
@@ -1387,7 +1394,8 @@ def llama_backend_free():
13871394 ...
13881395
13891396
1390- # //optional:
1397+ # // Optional: enable numa optimizations
1398+ # // TODO: deprecate and make part of llama_backend_init()
13911399# LLAMA_API void llama_numa_init(enum ggml_numa_strategy numa);
13921400@ctypes_function (
13931401 "llama_numa_init" ,
@@ -3192,6 +3200,207 @@ def llama_decode(ctx: llama_context_p, batch: llama_batch, /) -> int:
31923200 ...
31933201
31943202
3203+ # //
3204+ # // Extended batch API
3205+ # //
3206+
3207+ # struct llama_batch_ext;
3208+ llama_batch_ext_p = NewType ("llama_batch_ext_p" , int )
3209+ llama_batch_ext_p_ctypes = ctypes .c_void_p
3210+
3211+
3212+ # struct llama_embd {
3213+ # const float * data;
3214+ # size_t n_rows; // number of embedding rows in data
3215+ # size_t n_embd; // size of one row
3216+ # };
3217+ class llama_embd (ctypes .Structure ):
3218+ if TYPE_CHECKING :
3219+ data : CtypesPointer [ctypes .c_float ]
3220+ n_rows : int
3221+ n_embd : int
3222+
3223+ _fields_ = [
3224+ ("data" , ctypes .POINTER (ctypes .c_float )),
3225+ ("n_rows" , ctypes .c_size_t ),
3226+ ("n_embd" , ctypes .c_size_t ),
3227+ ]
3228+
3229+
3230+ # LLAMA_API struct llama_batch_ext * llama_batch_ext_init (struct llama_context * ctx);
3231+ @ctypes_function (
3232+ "llama_batch_ext_init" , [llama_context_p_ctypes ], llama_batch_ext_p_ctypes
3233+ )
3234+ def llama_batch_ext_init (ctx : llama_context_p , / ) -> Optional [llama_batch_ext_p ]: ...
3235+
3236+
3237+ # LLAMA_API void llama_batch_ext_free (struct llama_batch_ext * batch);
3238+ @ctypes_function ("llama_batch_ext_free" , [llama_batch_ext_p_ctypes ], None )
3239+ def llama_batch_ext_free (batch : llama_batch_ext_p , / ): ...
3240+
3241+
3242+ # LLAMA_API void llama_batch_ext_clear(struct llama_batch_ext * batch);
3243+ @ctypes_function ("llama_batch_ext_clear" , [llama_batch_ext_p_ctypes ], None )
3244+ def llama_batch_ext_clear (batch : llama_batch_ext_p , / ): ...
3245+
3246+
3247+ # // Add an input token to the batch, with default values:
3248+ # // id = LLAMA_TOKEN_NULL
3249+ # // embd = None
3250+ # // pos = not set, the caller must set it with llama_batch_ext_set_pos()
3251+ # // Returns the batch index (>= 0)
3252+ # // On error:
3253+ # // -1: batch is full
3254+ # // -2: token is invalid (id == LLAMA_TOKEN_NULL or invalid embd)
3255+ # // -3: invalid sequence id
3256+ # LLAMA_API int32_t llama_batch_ext_add(struct llama_batch_ext * batch, llama_seq_id seq_id);
3257+ @ctypes_function (
3258+ "llama_batch_ext_add" , [llama_batch_ext_p_ctypes , llama_seq_id ], ctypes .c_int32
3259+ )
3260+ def llama_batch_ext_add (batch : llama_batch_ext_p , seq_id : int , / ) -> int : ...
3261+
3262+
3263+ # // Add an input token to the batch, with a specified token ID or token embedding
3264+ # LLAMA_API int32_t llama_batch_ext_add_token(struct llama_batch_ext * batch, llama_seq_id seq_id, llama_token id);
3265+ @ctypes_function (
3266+ "llama_batch_ext_add_token" ,
3267+ [llama_batch_ext_p_ctypes , llama_seq_id , llama_token ],
3268+ ctypes .c_int32 ,
3269+ )
3270+ def llama_batch_ext_add_token (
3271+ batch : llama_batch_ext_p , seq_id : int , id : int , /
3272+ ) -> int : ...
3273+
3274+
3275+ # LLAMA_API int32_t llama_batch_ext_add_embd(struct llama_batch_ext * batch, llama_seq_id seq_id, struct llama_embd embd);
3276+ @ctypes_function (
3277+ "llama_batch_ext_add_embd" ,
3278+ [llama_batch_ext_p_ctypes , llama_seq_id , llama_embd ],
3279+ ctypes .c_int32 ,
3280+ )
3281+ def llama_batch_ext_add_embd (
3282+ batch : llama_batch_ext_p , seq_id : int , embd : llama_embd , /
3283+ ) -> int : ...
3284+
3285+
3286+ # // Add the token at index idx in the batch to another sequence id. The position will stays the same.
3287+ # // Note: this should be called before other _set() functions
3288+ # LLAMA_API bool llama_batch_ext_add_seq(
3289+ # struct llama_batch_ext * batch,
3290+ # int32_t idx,
3291+ # llama_seq_id seq_id);
3292+ @ctypes_function (
3293+ "llama_batch_ext_add_seq" ,
3294+ [llama_batch_ext_p_ctypes , ctypes .c_int32 , llama_seq_id ],
3295+ ctypes .c_bool ,
3296+ )
3297+ def llama_batch_ext_add_seq (
3298+ batch : llama_batch_ext_p , idx : int , seq_id : int , /
3299+ ) -> bool : ...
3300+
3301+
3302+ # // Set the token embedding for the token at index idx in the batch
3303+ # // use it after llama_batch_ext_add_token() to have an entry with both a token id and an embedding
3304+ # LLAMA_API bool llama_batch_ext_set_embd_token(
3305+ # struct llama_batch_ext * batch,
3306+ # int32_t idx,
3307+ # struct llama_embd embd);
3308+ @ctypes_function (
3309+ "llama_batch_ext_set_embd_token" ,
3310+ [llama_batch_ext_p_ctypes , ctypes .c_int32 , llama_embd ],
3311+ ctypes .c_bool ,
3312+ )
3313+ def llama_batch_ext_set_embd_token (
3314+ batch : llama_batch_ext_p , idx : int , embd : llama_embd , /
3315+ ) -> bool : ...
3316+
3317+
3318+ # // Set the "state" embedding for the token at index idx in the batch
3319+ # // "state" here means extra hidden state carried over from a previous stage, e.g.:
3320+ # // - MTP: state from N layers of the target model
3321+ # // - Qwen3 VL (deepstack): state from N layers of the vision encoder
3322+ # LLAMA_API bool llama_batch_ext_set_embd_state(
3323+ # struct llama_batch_ext * batch,
3324+ # int32_t idx,
3325+ # struct llama_embd embd);
3326+ @ctypes_function (
3327+ "llama_batch_ext_set_embd_state" ,
3328+ [llama_batch_ext_p_ctypes , ctypes .c_int32 , llama_embd ],
3329+ ctypes .c_bool ,
3330+ )
3331+ def llama_batch_ext_set_embd_state (
3332+ batch : llama_batch_ext_p , idx : int , embd : llama_embd , /
3333+ ) -> bool : ...
3334+
3335+
3336+ # // Set if output embeddings should be available for the token at index idx in the batch
3337+ # // Note: for now, this is equivalent to setting the output logits
3338+ # LLAMA_API bool llama_batch_ext_set_output_embd(
3339+ # struct llama_batch_ext * batch,
3340+ # int32_t idx,
3341+ # bool value);
3342+ @ctypes_function (
3343+ "llama_batch_ext_set_output_embd" ,
3344+ [llama_batch_ext_p_ctypes , ctypes .c_int32 , ctypes .c_bool ],
3345+ ctypes .c_bool ,
3346+ )
3347+ def llama_batch_ext_set_output_embd (
3348+ batch : llama_batch_ext_p , idx : int , value : bool , /
3349+ ) -> bool : ...
3350+
3351+
3352+ # // Set output logits for the token at index idx in the batch
3353+ # // Note: for now, this is equivalent to setting the output embd
3354+ # LLAMA_API bool llama_batch_ext_set_output_logits(
3355+ # struct llama_batch_ext * batch,
3356+ # int32_t idx,
3357+ # bool value);
3358+ @ctypes_function (
3359+ "llama_batch_ext_set_output_logits" ,
3360+ [llama_batch_ext_p_ctypes , ctypes .c_int32 , ctypes .c_bool ],
3361+ ctypes .c_bool ,
3362+ )
3363+ def llama_batch_ext_set_output_logits (
3364+ batch : llama_batch_ext_p , idx : int , value : bool , /
3365+ ) -> bool : ...
3366+
3367+
3368+ # // Set custom position for the token at index idx in the batch
3369+ # // For M-RoPE models:
3370+ # // - Embedding tokens must have multiple positions per token
3371+ # // - Text token only requires one single position per token
3372+ # LLAMA_API bool llama_batch_ext_set_pos(
3373+ # struct llama_batch_ext * batch,
3374+ # int32_t idx,
3375+ # const llama_pos * pos);
3376+ @ctypes_function (
3377+ "llama_batch_ext_set_pos" ,
3378+ [llama_batch_ext_p_ctypes , ctypes .c_int32 , ctypes .POINTER (llama_pos )],
3379+ ctypes .c_bool ,
3380+ )
3381+ def llama_batch_ext_set_pos (
3382+ batch : llama_batch_ext_p , idx : int , pos : CtypesPointerOrRef [llama_pos ], /
3383+ ) -> bool : ...
3384+
3385+
3386+ # // TODO: implement get_embeddings() and get_logits() for llama_batch_ext
3387+
3388+
3389+ # // Return values are the same as llama_decode()
3390+ # LLAMA_API int32_t llama_process(
3391+ # struct llama_context * ctx,
3392+ # enum llama_process_type type,
3393+ # struct llama_batch_ext * batch);
3394+ @ctypes_function (
3395+ "llama_process" ,
3396+ [llama_context_p_ctypes , ctypes .c_int , llama_batch_ext_p_ctypes ],
3397+ ctypes .c_int32 ,
3398+ )
3399+ def llama_process (ctx : llama_context_p , type : int , batch : llama_batch_ext_p , / ) -> int :
3400+ """Return values are the same as llama_decode()"""
3401+ ...
3402+
3403+
31953404# // Set the number of threads used for decoding
31963405# // n_threads is the number of threads used for generation (single token)
31973406# // n_threads_batch is the number of threads used for prompt and batch processing (multiple tokens)
@@ -3254,6 +3463,14 @@ def llama_set_causal_attn(ctx: llama_context_p, causal_attn: bool, /):
32543463 ...
32553464
32563465
3466+ # // Returns whether the context is currently using causal attention
3467+ # LLAMA_API bool llama_get_causal_attn(const struct llama_context * ctx);
3468+ @ctypes_function ("llama_get_causal_attn" , [llama_context_p_ctypes ], ctypes .c_bool )
3469+ def llama_get_causal_attn (ctx : llama_context_p , / ) -> bool :
3470+ """Returns whether the context is currently using causal attention"""
3471+ ...
3472+
3473+
32573474# // Set whether the model is in warmup mode or not
32583475# // If true, all model tensors are activated during llama_decode() to load and cache their weights.
32593476# LLAMA_API void llama_set_warmup(struct llama_context * ctx, bool warmup);
0 commit comments