2025-01-07 18:01:58 +01:00
#include "ggml.h"
#include "gguf.h"
2026-04-17 11:11:46 +03:00
#include "build-info.h"
2023-03-25 20:26:40 +02:00
#include "common.h"
2026-04-21 09:54:36 +03:00
#include "fit.h"
2024-09-15 20:46:12 +03:00
#include "log.h"
2023-08-28 17:59:39 +02:00
#include "llama.h"
2025-11-25 09:56:07 +08:00
#include "sampling.h"
2026-05-16 20:06:23 +08:00
#include "speculative.h"
2026-02-12 07:27:52 +01:00
#include "unicode.h"
2023-03-24 08:19:05 -07:00
2023-03-22 07:32:36 +02:00
#include <algorithm>
2024-06-04 21:23:39 +03:00
#include <cinttypes>
2024-10-12 08:21:51 +03:00
#include <climits>
2023-08-28 17:59:39 +02:00
#include <cmath>
2025-09-24 08:53:47 +02:00
#include <chrono>
2024-06-04 21:23:39 +03:00
#include <cstdarg>
2023-08-28 17:59:39 +02:00
#include <cstring>
#include <ctime>
2024-12-31 11:46:06 +11:00
#include <filesystem>
2023-08-28 17:59:39 +02:00
#include <fstream>
#include <iostream>
2024-06-04 21:23:39 +03:00
#include <iterator>
2023-06-06 21:33:23 +02:00
#include <regex>
2023-08-28 17:59:39 +02:00
#include <sstream>
#include <string>
2024-10-12 08:21:51 +03:00
#include <thread>
2023-08-28 17:59:39 +02:00
#include <unordered_set>
#include <vector>
2023-04-30 14:41:35 -04:00
#if defined(__APPLE__) && defined(__MACH__)
#include <sys/types.h>
#include <sys/sysctl.h>
#endif
2023-03-10 20:40:58 +02:00
2023-05-08 19:45:48 -07:00
#if defined(_WIN32)
#define WIN32_LEAN_AND_MEAN
2023-09-01 09:34:50 -04:00
#ifndef NOMINMAX
# define NOMINMAX
#endif
2023-08-28 17:59:39 +02:00
#include <locale>
2023-05-08 19:45:48 -07:00
#include <windows.h>
2025-08-14 03:03:57 -07:00
#include <string.h>
2023-04-08 17:49:39 +02:00
#include <fcntl.h>
#include <io.h>
2023-05-08 19:45:48 -07:00
#else
#include <sys/ioctl.h>
2023-08-28 17:59:39 +02:00
#include <sys/stat.h>
2023-05-08 19:45:48 -07:00
#include <unistd.h>
2023-03-28 17:09:55 +03:00
#endif
2023-03-12 17:15:00 -03:00
2025-09-27 02:03:33 +08:00
#if defined(__linux__)
#include <sys/types.h>
#include <pwd.h>
#endif
2026-07-07 03:05:20 +05:30
#if defined(_AIX)
#include <sys/systemcfg.h>
#endif
2023-06-16 21:23:53 +03:00
#if defined(_MSC_VER)
#pragma warning(disable: 4244 4267) // possible loss of data
#endif
2025-11-20 13:40:10 +02:00
common_time_meas :: common_time_meas ( int64_t & t_acc , bool disable ) : t_start_us ( disable ? - 1 : ggml_time_us ()), t_acc ( t_acc ) {}
common_time_meas ::~ common_time_meas () {
if ( t_start_us >= 0 ) {
t_acc += ggml_time_us () - t_start_us ;
}
}
2024-05-22 20:04:20 +03:00
//
// CPU utils
//
2026-04-28 09:07:33 +03:00
int32_t common_cpu_get_num_physical_cores () {
2026-07-07 03:05:20 +05:30
#if defined(_AIX)
int32_t logical_cpus = _system_configuration . ncpus ;
int32_t smt_threads = _system_configuration . smt_threads ;
if ( smt_threads > 0 ) {
return static_cast < int32_t > ( logical_cpus / smt_threads );
}
if ( logical_cpus > 0 ) {
return static_cast < int32_t > ( logical_cpus );
}
#elif defined(__linux__)
2023-05-14 22:25:42 -04:00
// enumerate the set of thread siblings, num entries is num cores
std :: unordered_set < std :: string > siblings ;
for ( uint32_t cpu = 0 ; cpu < UINT32_MAX ; ++ cpu ) {
2024-05-04 15:26:53 +02:00
std :: ifstream thread_siblings ( "/sys/devices/system/cpu/cpu"
2023-05-14 22:25:42 -04:00
+ std :: to_string ( cpu ) + "/topology/thread_siblings" );
if ( ! thread_siblings . is_open ()) {
break ; // no more cpus
2023-04-30 14:41:35 -04:00
}
2023-05-14 22:25:42 -04:00
std :: string line ;
if ( std :: getline ( thread_siblings , line )) {
siblings . insert ( line );
}
}
2023-09-07 13:22:29 -04:00
if ( ! siblings . empty ()) {
2023-05-14 22:25:42 -04:00
return static_cast < int32_t > ( siblings . size ());
2023-03-17 17:47:35 +00:00
}
2023-04-30 14:41:35 -04:00
#elif defined(__APPLE__) && defined(__MACH__)
int32_t num_physical_cores ;
size_t len = sizeof ( num_physical_cores );
int result = sysctlbyname ( "hw.perflevel0.physicalcpu" , & num_physical_cores , & len , NULL , 0 );
if ( result == 0 ) {
return num_physical_cores ;
}
result = sysctlbyname ( "hw.physicalcpu" , & num_physical_cores , & len , NULL , 0 );
if ( result == 0 ) {
return num_physical_cores ;
}
2024-08-16 14:23:12 +08:00
#elif defined(_WIN32) && (_WIN32_WINNT >= 0x0601) && !defined(__MINGW64__) // windows 7 and later
// TODO: windows + arm64 + mingw64
unsigned int n_threads_win = std :: thread :: hardware_concurrency ();
unsigned int default_threads = n_threads_win > 0 ? ( n_threads_win <= 4 ? n_threads_win : n_threads_win / 2 ) : 4 ;
DWORD buffer_size = 0 ;
if ( ! GetLogicalProcessorInformationEx ( RelationProcessorCore , nullptr , & buffer_size )) {
if ( GetLastError () != ERROR_INSUFFICIENT_BUFFER ) {
return default_threads ;
}
}
std :: vector < char > buffer ( buffer_size );
if ( ! GetLogicalProcessorInformationEx ( RelationProcessorCore , reinterpret_cast < PSYSTEM_LOGICAL_PROCESSOR_INFORMATION_EX > ( buffer . data ()), & buffer_size )) {
return default_threads ;
}
int32_t num_physical_cores = 0 ;
PSYSTEM_LOGICAL_PROCESSOR_INFORMATION_EX info = reinterpret_cast < PSYSTEM_LOGICAL_PROCESSOR_INFORMATION_EX > ( buffer . data ());
while ( buffer_size > 0 ) {
if ( info -> Relationship == RelationProcessorCore ) {
num_physical_cores += info -> Processor . GroupCount ;
}
buffer_size -= info -> Size ;
info = reinterpret_cast < PSYSTEM_LOGICAL_PROCESSOR_INFORMATION_EX > ( reinterpret_cast < char *> ( info ) + info -> Size );
}
return num_physical_cores > 0 ? num_physical_cores : default_threads ;
2023-04-30 14:41:35 -04:00
#endif
unsigned int n_threads = std :: thread :: hardware_concurrency ();
return n_threads > 0 ? ( n_threads <= 4 ? n_threads : n_threads / 2 ) : 4 ;
}
2023-03-17 17:47:35 +00:00
2024-04-20 13:27:12 +03:00
#if defined(__x86_64__) && defined(__linux__) && !defined(__ANDROID__)
2024-04-16 14:55:30 -04:00
#include <pthread.h>
static void cpuid ( unsigned leaf , unsigned subleaf ,
unsigned * eax , unsigned * ebx , unsigned * ecx , unsigned * edx ) {
__asm__ ( "movq \t %%rbx,%%rsi \n\t "
"cpuid \n\t "
"xchgq \t %%rbx,%%rsi"
: "=a" ( * eax ), "=S" ( * ebx ), "=c" ( * ecx ), "=d" ( * edx )
: "0" ( leaf ), "2" ( subleaf ));
}
static int pin_cpu ( int cpu ) {
cpu_set_t mask ;
CPU_ZERO ( & mask );
CPU_SET ( cpu , & mask );
return pthread_setaffinity_np ( pthread_self (), sizeof ( mask ), & mask );
}
static bool is_hybrid_cpu ( void ) {
unsigned eax , ebx , ecx , edx ;
cpuid ( 7 , 0 , & eax , & ebx , & ecx , & edx );
return !! ( edx & ( 1u << 15 ));
}
static bool is_running_on_efficiency_core ( void ) {
unsigned eax , ebx , ecx , edx ;
cpuid ( 0x1a , 0 , & eax , & ebx , & ecx , & edx );
int intel_atom = 0x20 ;
int core_type = ( eax & 0xff000000u ) >> 24 ;
return core_type == intel_atom ;
}
2024-05-22 20:04:20 +03:00
static int cpu_count_math_cpus ( int n_cpu ) {
2024-04-16 14:55:30 -04:00
int result = 0 ;
2024-05-22 20:04:20 +03:00
for ( int cpu = 0 ; cpu < n_cpu ; ++ cpu ) {
2024-04-16 14:55:30 -04:00
if ( pin_cpu ( cpu )) {
return - 1 ;
}
if ( is_running_on_efficiency_core ()) {
continue ; // efficiency cores harm lockstep threading
}
++ cpu ; // hyperthreading isn't useful for linear algebra
++ result ;
}
return result ;
}
#endif // __x86_64__ && __linux__
/**
* Returns number of CPUs on system that are useful for math.
*/
2026-04-28 09:07:33 +03:00
int32_t common_cpu_get_num_math () {
2024-04-20 13:27:12 +03:00
#if defined(__x86_64__) && defined(__linux__) && !defined(__ANDROID__)
2024-05-22 20:04:20 +03:00
int n_cpu = sysconf ( _SC_NPROCESSORS_ONLN );
if ( n_cpu < 1 ) {
2026-04-28 09:07:33 +03:00
return common_cpu_get_num_physical_cores ();
2024-04-16 14:55:30 -04:00
}
if ( is_hybrid_cpu ()) {
cpu_set_t affinity ;
if ( ! pthread_getaffinity_np ( pthread_self (), sizeof ( affinity ), & affinity )) {
2024-05-22 20:04:20 +03:00
int result = cpu_count_math_cpus ( n_cpu );
2024-04-16 14:55:30 -04:00
pthread_setaffinity_np ( pthread_self (), sizeof ( affinity ), & affinity );
if ( result > 0 ) {
return result ;
}
}
}
2026-07-07 03:05:20 +05:30
#elif defined(__powerpc64__) || defined(__powerpc__)
int32_t smt_factor = 1 ;
int phy_cpus = common_cpu_get_num_physical_cores ();
int logical_cpus = sysconf ( _SC_NPROCESSORS_ONLN );
if ( phy_cpus > 0 && logical_cpus > phy_cpus ) {
smt_factor = logical_cpus / phy_cpus ;
}
return phy_cpus * std :: min ( smt_factor , 2 );
2024-04-16 14:55:30 -04:00
#endif
2026-04-28 09:07:33 +03:00
return common_cpu_get_num_physical_cores ();
2024-04-16 14:55:30 -04:00
}
2024-08-29 19:20:53 -04:00
// Helper for setting process priority
#if defined(_WIN32)
bool set_process_priority ( enum ggml_sched_priority prio ) {
if ( prio == GGML_SCHED_PRIO_NORMAL ) {
return true ;
}
DWORD p = NORMAL_PRIORITY_CLASS ;
switch ( prio ) {
2025-05-31 15:39:19 -07:00
case GGML_SCHED_PRIO_LOW : p = BELOW_NORMAL_PRIORITY_CLASS ; break ;
2024-08-29 19:20:53 -04:00
case GGML_SCHED_PRIO_NORMAL : p = NORMAL_PRIORITY_CLASS ; break ;
case GGML_SCHED_PRIO_MEDIUM : p = ABOVE_NORMAL_PRIORITY_CLASS ; break ;
case GGML_SCHED_PRIO_HIGH : p = HIGH_PRIORITY_CLASS ; break ;
case GGML_SCHED_PRIO_REALTIME : p = REALTIME_PRIORITY_CLASS ; break ;
}
if ( ! SetPriorityClass ( GetCurrentProcess (), p )) {
2026-06-28 08:52:15 +03:00
COM_WRN ( "failed to set process priority class %d : (%d) \n " , prio , ( int ) GetLastError ());
2024-08-29 19:20:53 -04:00
return false ;
}
return true ;
}
#else // MacOS and POSIX
#include <sys/types.h>
#include <sys/resource.h>
bool set_process_priority ( enum ggml_sched_priority prio ) {
if ( prio == GGML_SCHED_PRIO_NORMAL ) {
return true ;
}
int p = 0 ;
switch ( prio ) {
2025-05-31 15:39:19 -07:00
case GGML_SCHED_PRIO_LOW : p = 5 ; break ;
2024-08-29 19:20:53 -04:00
case GGML_SCHED_PRIO_NORMAL : p = 0 ; break ;
case GGML_SCHED_PRIO_MEDIUM : p = - 5 ; break ;
case GGML_SCHED_PRIO_HIGH : p = - 10 ; break ;
case GGML_SCHED_PRIO_REALTIME : p = - 20 ; break ;
}
2025-12-29 17:07:49 +08:00
if ( setpriority ( PRIO_PROCESS , 0 , p ) != 0 ) {
2026-06-28 08:52:15 +03:00
COM_WRN ( "failed to set process priority %d : %s (%d) \n " , prio , strerror ( errno ), errno );
2024-08-29 19:20:53 -04:00
return false ;
}
return true ;
}
#endif
2024-05-22 20:04:20 +03:00
//
// CLI argument parsing
//
2023-05-02 18:46:20 -07:00
2024-05-22 20:04:20 +03:00
2026-04-28 09:07:33 +03:00
void postprocess_cpu_params ( common_cpu_params & cpuparams , const common_cpu_params * role_model ) {
2024-08-29 19:20:53 -04:00
int32_t n_set = 0 ;
if ( cpuparams . n_threads < 0 ) {
// Assuming everything about cpuparams is invalid
if ( role_model != nullptr ) {
cpuparams = * role_model ;
} else {
2026-04-28 09:07:33 +03:00
cpuparams . n_threads = common_cpu_get_num_math ();
2024-08-29 19:20:53 -04:00
}
}
for ( int32_t i = 0 ; i < GGML_MAX_N_THREADS ; i ++ ) {
if ( cpuparams . cpumask [ i ]) {
n_set ++ ;
}
}
if ( n_set && n_set < cpuparams . n_threads ) {
// Not enough set bits, may experience performance issues.
2026-06-28 08:52:15 +03:00
COM_WRN ( "Not enough set bits in CPU mask (%d) to satisfy requested thread count: %d \n " , n_set , cpuparams . n_threads );
2024-08-29 19:20:53 -04:00
}
}
bool parse_cpu_range ( const std :: string & range , bool ( & boolmask )[ GGML_MAX_N_THREADS ]) {
size_t dash_loc = range . find ( '-' );
if ( dash_loc == std :: string :: npos ) {
2026-06-28 08:52:15 +03:00
COM_ERR ( "%s" , "Format of CPU range is invalid! Expected [<start>]-[<end>]. \n " );
2024-08-29 19:20:53 -04:00
return false ;
}
size_t start_i ;
size_t end_i ;
if ( dash_loc == 0 ) {
start_i = 0 ;
} else {
start_i = std :: stoull ( range . substr ( 0 , dash_loc ));
if ( start_i >= GGML_MAX_N_THREADS ) {
2026-06-28 08:52:15 +03:00
COM_ERR ( "%s" , "Start index out of bounds! \n " );
2024-08-29 19:20:53 -04:00
return false ;
}
}
if ( dash_loc == range . length () - 1 ) {
end_i = GGML_MAX_N_THREADS - 1 ;
} else {
end_i = std :: stoull ( range . substr ( dash_loc + 1 ));
if ( end_i >= GGML_MAX_N_THREADS ) {
2026-06-28 08:52:15 +03:00
COM_ERR ( "%s" , "End index out of bounds! \n " );
2024-08-29 19:20:53 -04:00
return false ;
}
}
for ( size_t i = start_i ; i <= end_i ; i ++ ) {
boolmask [ i ] = true ;
}
return true ;
}
bool parse_cpu_mask ( const std :: string & mask , bool ( & boolmask )[ GGML_MAX_N_THREADS ]) {
// Discard potential 0x prefix
size_t start_i = 0 ;
if ( mask . length () >= 2 && mask . substr ( 0 , 2 ) == "0x" ) {
start_i = 2 ;
}
size_t num_digits = mask . length () - start_i ;
2026-06-28 08:52:15 +03:00
num_digits = std :: min < size_t > ( num_digits , 128 );
2024-08-29 19:20:53 -04:00
size_t end_i = num_digits + start_i ;
for ( size_t i = start_i , n = ( num_digits * 4 - 1 ); i < end_i ; i ++ , n -= 4 ) {
char c = mask . at ( i );
int8_t id = c ;
if (( c >= '0' && c <= '9' )) {
id -= '0' ;
} else if ( c >= 'a' && c <= 'f' ) {
id -= 'a' - 10 ;
} else if ( c >= 'A' && c <= 'F' ) {
id -= 'A' - 10 ;
} else {
2026-06-28 08:52:15 +03:00
COM_ERR ( "Invalid hex character '%c' at position %d \n " , c , int32_t ( i ));
2024-08-29 19:20:53 -04:00
return false ;
}
boolmask [ n ] = boolmask [ n ] || (( id & 8 ) != 0 );
boolmask [ n - 1 ] = boolmask [ n - 1 ] || (( id & 4 ) != 0 );
boolmask [ n - 2 ] = boolmask [ n - 2 ] || (( id & 2 ) != 0 );
boolmask [ n - 3 ] = boolmask [ n - 3 ] || (( id & 1 ) != 0 );
}
return true ;
}
2024-10-10 22:57:42 +02:00
void common_init () {
2026-03-31 12:53:41 +02:00
#if defined(_WIN32)
SetConsoleOutputCP ( CP_UTF8 );
SetConsoleCP ( CP_UTF8 );
#endif
2026-05-14 13:05:52 +03:00
common_log_set_prefix ( common_log_main (), true );
common_log_set_timestamps ( common_log_main (), true );
2024-09-15 20:46:12 +03:00
2026-05-14 13:05:52 +03:00
llama_log_set ( common_log_default_callback , NULL );
}
2026-05-16 21:21:06 +02:00
void common_params_print_info ( const common_params & params , bool print_devices ) {
2024-09-15 20:46:12 +03:00
#ifdef NDEBUG
const char * build_type = "" ;
#else
const char * build_type = " (debug)" ;
#endif
2026-06-28 08:52:15 +03:00
COM_TRC ( "%s: build %d (%s) with %s for %s%s \n " , __func__ , llama_build_number (), llama_commit (), llama_compiler (), llama_build_target (), build_type );
2024-09-15 20:46:12 +03:00
2026-06-28 08:52:15 +03:00
COM_INF ( "%s: verbosity = %d (adjust with the `-lv N` CLI arg) \n " , __func__ , common_log_get_verbosity_thold ());
2026-05-16 21:21:06 +02:00
// device enumeration creates a primary context on CUDA backends, skip it when the caller does not own any device
if ( print_devices ) {
2026-06-28 08:52:15 +03:00
COM_TRC ( "%s" , "device_info: \n " );
2026-05-16 21:21:06 +02:00
for ( size_t i = 0 ; i < ggml_backend_dev_count (); ++ i ) {
auto * dev = ggml_backend_dev_get ( i );
size_t free , total ;
ggml_backend_dev_memory ( dev , & free , & total );
2026-06-28 08:52:15 +03:00
COM_TRC ( " - %-8s: %s (%zu MiB, %zu MiB free) \n " , ggml_backend_dev_name ( dev ), ggml_backend_dev_description ( dev ), total / 1024 / 1024 , free / 1024 / 1024 );
2026-05-16 21:21:06 +02:00
}
2026-05-14 13:05:52 +03:00
}
2026-06-28 08:52:15 +03:00
COM_TRC ( "%s \n " , common_params_get_system_info ( params ). c_str ());
2024-09-15 20:46:12 +03:00
}
2024-10-10 22:57:42 +02:00
std :: string common_params_get_system_info ( const common_params & params ) {
2023-09-28 21:42:38 +02:00
std :: ostringstream os ;
2024-08-29 19:20:53 -04:00
os << "system_info: n_threads = " << params . cpuparams . n_threads ;
if ( params . cpuparams_batch . n_threads != - 1 ) {
os << " (n_threads_batch = " << params . cpuparams_batch . n_threads << ")" ;
2023-09-28 21:42:38 +02:00
}
2024-08-16 14:23:12 +08:00
#if defined(_WIN32) && (_WIN32_WINNT >= 0x0601) && !defined(__MINGW64__) // windows 7 and later
// TODO: windows + arm64 + mingw64
DWORD logicalProcessorCount = GetActiveProcessorCount ( ALL_PROCESSOR_GROUPS );
os << " / " << logicalProcessorCount << " | " << llama_print_system_info ();
#else
2023-09-28 21:42:38 +02:00
os << " / " << std :: thread :: hardware_concurrency () << " | " << llama_print_system_info ();
2024-08-16 14:23:12 +08:00
#endif
2023-09-28 21:42:38 +02:00
return os . str ();
}
2024-05-22 20:04:20 +03:00
//
// String utils
//
2024-10-12 08:21:51 +03:00
std :: string string_format ( const char * fmt , ...) {
va_list ap ;
va_list ap2 ;
va_start ( ap , fmt );
va_copy ( ap2 , ap );
int size = vsnprintf ( NULL , 0 , fmt , ap );
GGML_ASSERT ( size >= 0 && size < INT_MAX ); // NOLINT
std :: vector < char > buf ( size + 1 );
int size2 = vsnprintf ( buf . data (), size + 1 , fmt , ap2 );
GGML_ASSERT ( size2 == size );
va_end ( ap2 );
va_end ( ap );
return std :: string ( buf . data (), size );
}
2024-05-22 20:04:20 +03:00
std :: string string_strip ( const std :: string & str ) {
size_t start = 0 ;
size_t end = str . size ();
while ( start < end && std :: isspace ( str [ start ])) {
start ++ ;
}
while ( end > start && std :: isspace ( str [ end - 1 ])) {
end -- ;
}
return str . substr ( start , end - start );
}
2026-05-25 07:56:18 +02:00
std :: string string_lcs ( std :: string_view a , std :: string_view b ) {
if ( a . empty () || b . empty ()) return {};
std :: vector < std :: vector < size_t >> dp ( a . size () + 1 , std :: vector < size_t > ( b . size () + 1 , 0 ));
size_t best_len = 0 ;
size_t best_end_a = 0 ;
for ( size_t i = 1 ; i <= a . size (); ++ i ) {
for ( size_t j = 1 ; j <= b . size (); ++ j ) {
if ( a [ i - 1 ] == b [ j - 1 ]) {
dp [ i ][ j ] = dp [ i - 1 ][ j - 1 ] + 1 ;
if ( dp [ i ][ j ] > best_len ) {
best_len = dp [ i ][ j ];
best_end_a = i ;
}
}
}
}
return std :: string ( a . substr ( best_end_a - best_len , best_len ));
}
2024-05-22 20:04:20 +03:00
std :: string string_get_sortable_timestamp () {
using clock = std :: chrono :: system_clock ;
const clock :: time_point current_time = clock :: now ();
const time_t as_time_t = clock :: to_time_t ( current_time );
char timestamp_no_ns [ 100 ];
std :: strftime ( timestamp_no_ns , 100 , "%Y_%m_%d-%H_%M_%S" , std :: localtime ( & as_time_t ));
const int64_t ns = std :: chrono :: duration_cast < std :: chrono :: nanoseconds > (
current_time . time_since_epoch () % 1000000000 ). count ();
char timestamp_ns [ 11 ];
snprintf ( timestamp_ns , 11 , "%09" PRId64 , ns );
return std :: string ( timestamp_no_ns ) + "." + std :: string ( timestamp_ns );
}
2024-08-09 18:23:52 +03:00
void string_replace_all ( std :: string & s , const std :: string & search , const std :: string & replace ) {
if ( search . empty ()) {
2024-08-25 23:09:53 -07:00
return ;
2024-08-09 18:23:52 +03:00
}
2024-08-25 23:09:53 -07:00
std :: string builder ;
builder . reserve ( s . length ());
2024-08-09 18:23:52 +03:00
size_t pos = 0 ;
2024-08-25 23:09:53 -07:00
size_t last_pos = 0 ;
while (( pos = s . find ( search , last_pos )) != std :: string :: npos ) {
builder . append ( s , last_pos , pos - last_pos );
builder . append ( replace );
last_pos = pos + search . length ();
2024-08-09 18:23:52 +03:00
}
2024-08-25 23:09:53 -07:00
builder . append ( s , last_pos , std :: string :: npos );
s = std :: move ( builder );
2024-08-09 18:23:52 +03:00
}
2025-03-05 13:05:13 +00:00
std :: string regex_escape ( const std :: string & s ) {
static const std :: regex special_chars ( "[.^$|()*+? \\ [ \\ ]{} \\\\ ]" );
2025-06-11 17:19:44 -03:00
return std :: regex_replace ( s , special_chars , " \\ $&" );
2025-03-05 13:05:13 +00:00
}
2025-01-22 09:51:44 +00:00
std :: string string_join ( const std :: vector < std :: string > & values , const std :: string & separator ) {
std :: ostringstream result ;
for ( size_t i = 0 ; i < values . size (); ++ i ) {
if ( i > 0 ) {
result << separator ;
}
result << values [ i ];
}
return result . str ();
}
std :: vector < std :: string > string_split ( const std :: string & str , const std :: string & delimiter ) {
std :: vector < std :: string > parts ;
size_t start = 0 ;
size_t end = str . find ( delimiter );
while ( end != std :: string :: npos ) {
parts . push_back ( str . substr ( start , end - start ));
start = end + delimiter . length ();
end = str . find ( delimiter , start );
}
parts . push_back ( str . substr ( start ));
return parts ;
}
std :: string string_repeat ( const std :: string & str , size_t n ) {
if ( n == 0 ) {
return "" ;
}
std :: string result ;
result . reserve ( str . length () * n );
for ( size_t i = 0 ; i < n ; ++ i ) {
result += str ;
}
return result ;
}
2024-09-15 20:46:12 +03:00
std :: string string_from ( bool value ) {
return value ? "true" : "false" ;
}
std :: string string_from ( const std :: vector < int > & values ) {
std :: stringstream buf ;
buf << "[ " ;
bool first = true ;
for ( auto e : values ) {
if ( first ) {
first = false ;
} else {
buf << ", " ;
}
buf << std :: to_string ( e );
}
buf << " ]" ;
return buf . str ();
}
std :: string string_from ( const struct llama_context * ctx , const std :: vector < llama_token > & tokens ) {
std :: stringstream buf ;
buf << "[ " ;
bool first = true ;
for ( const auto & token : tokens ) {
if ( ! first ) {
buf << ", " ;
} else {
first = false ;
}
2024-10-10 22:57:42 +02:00
auto detokenized = common_token_to_piece ( ctx , token );
2024-09-15 20:46:12 +03:00
buf << "'" << detokenized << "'"
<< ":" << std :: to_string ( token );
}
buf << " ]" ;
return buf . str ();
}
std :: string string_from ( const struct llama_context * ctx , const struct llama_batch & batch ) {
std :: stringstream buf ;
buf << "[ " ;
bool first = true ;
for ( int i = 0 ; i < batch . n_tokens ; ++ i ) {
if ( ! first ) {
buf << ", " ;
} else {
first = false ;
}
2024-10-10 22:57:42 +02:00
auto detokenized = common_token_to_piece ( ctx , batch . token [ i ]);
2024-09-15 20:46:12 +03:00
2024-11-25 09:58:41 +02:00
buf << " \n " << std :: to_string ( i )
<< ", token '" << detokenized << "'"
<< ", pos " << std :: to_string ( batch . pos [ i ])
<< ", n_seq_id " << std :: to_string ( batch . n_seq_id [ i ])
<< ", seq_id " << std :: to_string ( batch . seq_id [ i ][ 0 ])
<< ", logits " << std :: to_string ( batch . logits [ i ]);
2024-09-15 20:46:12 +03:00
}
buf << " ]" ;
return buf . str ();
}
2024-05-22 20:04:20 +03:00
void string_process_escapes ( std :: string & input ) {
std :: size_t input_len = input . length ();
std :: size_t output_idx = 0 ;
for ( std :: size_t input_idx = 0 ; input_idx < input_len ; ++ input_idx ) {
if ( input [ input_idx ] == '\\' && input_idx + 1 < input_len ) {
switch ( input [ ++ input_idx ]) {
case 'n' : input [ output_idx ++ ] = '\n' ; break ;
case 'r' : input [ output_idx ++ ] = '\r' ; break ;
case 't' : input [ output_idx ++ ] = '\t' ; break ;
case '\'' : input [ output_idx ++ ] = '\'' ; break ;
case '\"' : input [ output_idx ++ ] = '\"' ; break ;
case '\\' : input [ output_idx ++ ] = '\\' ; break ;
case 'x' :
// Handle \x12, etc
if ( input_idx + 2 < input_len ) {
const char x [ 3 ] = { input [ input_idx + 1 ], input [ input_idx + 2 ], 0 };
char * err_p = nullptr ;
const long val = std :: strtol ( x , & err_p , 16 );
if ( err_p == x + 2 ) {
input_idx += 2 ;
input [ output_idx ++ ] = char ( val );
break ;
}
}
// fall through
default : input [ output_idx ++ ] = '\\' ;
input [ output_idx ++ ] = input [ input_idx ]; break ;
}
} else {
input [ output_idx ++ ] = input [ input_idx ];
}
}
input . resize ( output_idx );
}
bool string_parse_kv_override ( const char * data , std :: vector < llama_model_kv_override > & overrides ) {
const char * sep = strchr ( data , '=' );
if ( sep == nullptr || sep - data >= 128 ) {
2026-06-28 08:52:15 +03:00
COM_ERR ( "%s: malformed KV override '%s' \n " , __func__ , data );
2024-05-22 20:04:20 +03:00
return false ;
}
llama_model_kv_override kvo ;
std :: strncpy ( kvo . key , data , sep - data );
kvo . key [ sep - data ] = 0 ;
sep ++ ;
if ( strncmp ( sep , "int:" , 4 ) == 0 ) {
sep += 4 ;
kvo . tag = LLAMA_KV_OVERRIDE_TYPE_INT ;
kvo . val_i64 = std :: atol ( sep );
} else if ( strncmp ( sep , "float:" , 6 ) == 0 ) {
sep += 6 ;
kvo . tag = LLAMA_KV_OVERRIDE_TYPE_FLOAT ;
kvo . val_f64 = std :: atof ( sep );
} else if ( strncmp ( sep , "bool:" , 5 ) == 0 ) {
sep += 5 ;
kvo . tag = LLAMA_KV_OVERRIDE_TYPE_BOOL ;
if ( std :: strcmp ( sep , "true" ) == 0 ) {
kvo . val_bool = true ;
} else if ( std :: strcmp ( sep , "false" ) == 0 ) {
kvo . val_bool = false ;
} else {
2026-06-28 08:52:15 +03:00
COM_ERR ( "%s: invalid boolean value for KV override '%s' \n " , __func__ , data );
2024-05-22 20:04:20 +03:00
return false ;
}
} else if ( strncmp ( sep , "str:" , 4 ) == 0 ) {
sep += 4 ;
kvo . tag = LLAMA_KV_OVERRIDE_TYPE_STR ;
if ( strlen ( sep ) > 127 ) {
2026-06-28 08:52:15 +03:00
COM_ERR ( "%s: malformed KV override '%s', value cannot exceed 127 chars \n " , __func__ , data );
2024-05-22 20:04:20 +03:00
return false ;
}
strncpy ( kvo . val_str , sep , 127 );
kvo . val_str [ 127 ] = '\0' ;
} else {
2026-06-28 08:52:15 +03:00
COM_ERR ( "%s: invalid type for KV override '%s' \n " , __func__ , data );
2024-05-22 20:04:20 +03:00
return false ;
}
overrides . emplace_back ( std :: move ( kvo ));
return true ;
}
2026-03-28 19:57:37 +01:00
static inline bool glob_class_match ( const char c , const char * pattern , const char * class_end ) {
const char * class_start = pattern ;
bool negated = false ;
if ( * class_start == '!' ) {
negated = true ;
class_start ++ ;
}
// If first character after negation is ']' or '-', treat it as literal
if ( * class_start == ']' || * class_start == '-' ) {
if ( class_start < class_end && * class_start == c ) {
return ! negated ;
}
class_start ++ ;
}
bool matched = false ;
while ( class_start < class_end ) {
if ( class_start + 2 < class_end && class_start [ 1 ] == '-' && class_start [ 2 ] != ']' ) {
char start_char = * class_start ;
char end_char = class_start [ 2 ];
if ( c >= start_char && c <= end_char ) {
matched = true ;
break ;
}
class_start += 3 ;
} else {
if ( * class_start == c ) {
matched = true ;
break ;
}
class_start ++ ;
}
}
return negated ? ! matched : matched ;
}
// simple glob: * matches non-/ chars, ** matches anything including /, [] matches character class
2026-03-28 02:33:04 +01:00
static inline bool glob_match ( const char * pattern , const char * str ) {
if ( * pattern == '\0' ) {
return * str == '\0' ;
}
if ( pattern [ 0 ] == '*' && pattern [ 1 ] == '*' ) {
const char * p = pattern + 2 ;
if ( glob_match ( p , str )) return true ;
if ( * str != '\0' ) return glob_match ( pattern , str + 1 );
return false ;
}
if ( * pattern == '*' ) {
const char * p = pattern + 1 ;
for (; * str != '\0' && * str != '/' ; str ++ ) {
if ( glob_match ( p , str )) return true ;
}
return glob_match ( p , str );
}
if ( * pattern == '?' && * str != '\0' && * str != '/' ) {
return glob_match ( pattern + 1 , str + 1 );
}
2026-03-28 19:57:37 +01:00
if ( * pattern == '[' ) {
const char * class_end = pattern + 1 ;
// If first character after '[' is ']' or '-', treat it as literal
if ( * class_end == ']' || * class_end == '-' ) {
class_end ++ ;
}
while ( * class_end != '\0' && * class_end != ']' ) {
class_end ++ ;
}
if ( * class_end == ']' ) {
if ( * str == '\0' ) return false ;
bool matched = glob_class_match ( * str , pattern + 1 , class_end );
return matched && glob_match ( class_end + 1 , str + 1 );
} else {
if ( * str == '[' ) {
return glob_match ( pattern + 1 , str + 1 );
}
return false ;
}
}
2026-03-28 02:33:04 +01:00
if ( * pattern == * str ) {
return glob_match ( pattern + 1 , str + 1 );
}
return false ;
}
bool glob_match ( const std :: string & pattern , const std :: string & str ) {
return glob_match ( pattern . c_str (), str . c_str ());
}
2024-05-22 20:04:20 +03:00
//
// Filesystem utils
//
2024-04-08 20:43:30 +08:00
// Validate if a filename is safe to use
// To validate a full path, split the path by the OS-specific path separator, and validate each part with this function
2025-12-02 22:49:20 +01:00
bool fs_validate_filename ( const std :: string & filename , bool allow_subdirs ) {
2024-04-08 20:43:30 +08:00
if ( ! filename . length ()) {
// Empty filename invalid
return false ;
}
if ( filename . length () > 255 ) {
// Limit at common largest possible filename on Linux filesystems
// to avoid unnecessary further validation
// (On systems with smaller limits it will be caught by the OS)
return false ;
}
2026-02-12 07:27:52 +01:00
size_t offset = 0 ;
while ( offset < filename . size ()) {
2026-03-06 21:01:00 +01:00
utf8_parse_result result = common_parse_utf8_codepoint ( filename , offset );
2025-06-19 20:49:48 +08:00
2026-02-12 07:27:52 +01:00
if ( result . status != utf8_parse_result :: SUCCESS ) {
2024-04-08 20:43:30 +08:00
return false ;
}
2026-02-12 07:27:52 +01:00
uint32_t c = result . codepoint ;
2024-04-08 20:43:30 +08:00
2026-02-12 07:27:52 +01:00
if (( result . bytes_consumed == 2 && c < 0x80 ) ||
( result . bytes_consumed == 3 && c < 0x800 ) ||
( result . bytes_consumed == 4 && c < 0x10000 )) {
return false ;
}
// Check for forbidden codepoints:
// - Control characters
// - Unicode equivalents of illegal characters
// - UTF-16 surrogate pairs
// - UTF-8 replacement character
// - Byte order mark (BOM)
// - Illegal characters: / \ : * ? " < > |
2024-04-08 20:43:30 +08:00
if ( c <= 0x1F // Control characters (C0)
|| c == 0x7F // Control characters (DEL)
|| ( c >= 0x80 && c <= 0x9F ) // Control characters (C1)
|| c == 0xFF0E // Fullwidth Full Stop (period equivalent)
|| c == 0x2215 // Division Slash (forward slash equivalent)
|| c == 0x2216 // Set Minus (backslash equivalent)
|| ( c >= 0xD800 && c <= 0xDFFF ) // UTF-16 surrogate pairs
2026-02-12 07:27:52 +01:00
|| c > 0x10FFFF // Max Unicode limit
2024-04-08 20:43:30 +08:00
|| c == 0xFFFD // Replacement Character (UTF-8)
|| c == 0xFEFF // Byte Order Mark (BOM)
2025-12-02 22:49:20 +01:00
|| c == ':' || c == '*' // Illegal characters
2024-04-08 20:43:30 +08:00
|| c == '?' || c == '"' || c == '<' || c == '>' || c == '|' ) {
return false ;
}
2025-12-02 22:49:20 +01:00
if ( ! allow_subdirs && ( c == '/' || c == '\\' )) {
// Subdirectories not allowed, reject path separators
return false ;
}
2026-02-12 07:27:52 +01:00
offset += result . bytes_consumed ;
2024-04-08 20:43:30 +08:00
}
// Reject any leading or trailing ' ', or any trailing '.', these are stripped on Windows and will cause a different filename
// Unicode and other whitespace is not affected, only 0x20 space
if ( filename . front () == ' ' || filename . back () == ' ' || filename . back () == '.' ) {
return false ;
}
// Reject any ".." (currently stricter than necessary, it should be fine to just check for == ".." instead)
if ( filename . find ( ".." ) != std :: string :: npos ) {
return false ;
}
// Reject "."
if ( filename == "." ) {
return false ;
}
return true ;
}
2025-06-16 08:11:43 -07:00
#include <iostream>
2025-12-04 11:06:49 +01:00
#ifdef _WIN32
static std :: wstring utf8_to_wstring ( const std :: string & str ) {
if ( str . empty ()) {
return std :: wstring ();
}
int size = MultiByteToWideChar ( CP_UTF8 , 0 , str . c_str (), ( int ) str . size (), NULL , 0 );
if ( size <= 0 ) {
return std :: wstring ();
}
std :: wstring wstr ( size , 0 );
MultiByteToWideChar ( CP_UTF8 , 0 , str . c_str (), ( int ) str . size (), & wstr [ 0 ], size );
return wstr ;
}
#endif
2024-05-22 20:04:20 +03:00
// returns true if successful, false otherwise
bool fs_create_directory_with_parents ( const std :: string & path ) {
#ifdef _WIN32
2025-12-04 11:06:49 +01:00
std :: wstring wpath = utf8_to_wstring ( path );
2023-12-05 15:05:51 +05:00
2024-05-22 20:04:20 +03:00
// if the path already exists, check whether it's a directory
const DWORD attributes = GetFileAttributesW ( wpath . c_str ());
if (( attributes != INVALID_FILE_ATTRIBUTES ) && ( attributes & FILE_ATTRIBUTE_DIRECTORY )) {
return true ;
2024-02-11 13:43:31 +00:00
}
2024-05-22 20:04:20 +03:00
size_t pos_slash = 0 ;
2024-04-29 16:58:41 +03:00
2024-05-22 20:04:20 +03:00
// process path from front to back, procedurally creating directories
while (( pos_slash = path . find ( '\\' , pos_slash )) != std :: string :: npos ) {
const std :: wstring subpath = wpath . substr ( 0 , pos_slash );
2024-02-16 11:33:25 +00:00
2025-06-16 08:11:43 -07:00
pos_slash += 1 ;
// skip the drive letter, in some systems it can return an access denied error
if ( subpath . length () == 2 && subpath [ 1 ] == ':' ) {
continue ;
}
const bool success = CreateDirectoryW ( subpath . c_str (), NULL );
2024-05-22 20:04:20 +03:00
if ( ! success ) {
const DWORD error = GetLastError ();
2023-12-05 15:05:51 +05:00
2024-05-22 20:04:20 +03:00
// if the path already exists, ensure that it's a directory
if ( error == ERROR_ALREADY_EXISTS ) {
const DWORD attributes = GetFileAttributesW ( subpath . c_str ());
if ( attributes == INVALID_FILE_ATTRIBUTES || ! ( attributes & FILE_ATTRIBUTE_DIRECTORY )) {
return false ;
2024-02-16 11:33:25 +00:00
}
2024-05-22 20:04:20 +03:00
} else {
return false ;
2024-02-16 11:33:25 +00:00
}
}
2023-12-05 15:05:51 +05:00
}
2024-02-11 13:43:31 +00:00
2024-05-22 20:04:20 +03:00
return true ;
#else
// if the path already exists, check whether it's a directory
struct stat info ;
if ( stat ( path . c_str (), & info ) == 0 ) {
return S_ISDIR ( info . st_mode );
}
2024-02-11 13:43:31 +00:00
2024-05-22 20:04:20 +03:00
size_t pos_slash = 1 ; // skip leading slashes for directory creation
// process path from front to back, procedurally creating directories
while (( pos_slash = path . find ( '/' , pos_slash )) != std :: string :: npos ) {
const std :: string subpath = path . substr ( 0 , pos_slash );
struct stat info ;
// if the path already exists, ensure that it's a directory
if ( stat ( subpath . c_str (), & info ) == 0 ) {
if ( ! S_ISDIR ( info . st_mode )) {
return false ;
}
} else {
// create parent directories
const int ret = mkdir ( subpath . c_str (), 0755 );
if ( ret != 0 ) {
return false ;
}
2024-02-11 13:43:31 +00:00
}
2024-05-22 20:04:20 +03:00
pos_slash += 1 ;
2024-02-11 13:43:31 +00:00
}
2024-05-22 20:04:20 +03:00
return true ;
#endif // _WIN32
2024-02-11 13:43:31 +00:00
}
2025-12-02 22:49:20 +01:00
bool fs_is_directory ( const std :: string & path ) {
std :: filesystem :: path dir ( path );
return std :: filesystem :: exists ( dir ) && std :: filesystem :: is_directory ( dir );
}
2026-08-03 18:58:15 +02:00
std :: string common_get_env ( const std :: string & name ) {
const char * value = std :: getenv ( name . c_str ());
return value == nullptr ? "" : value ;
}
void common_set_env ( const std :: string & name , const std :: string & value ) {
#if defined(_WIN32)
_putenv_s ( name . c_str (), value . c_str ());
#else
if ( value . empty ()) {
unsetenv ( name . c_str ());
} else {
setenv ( name . c_str (), value . c_str (), 1 );
}
#endif
}
2024-05-22 20:04:20 +03:00
std :: string fs_get_cache_directory () {
std :: string cache_directory = "" ;
2024-05-25 05:30:59 +02:00
auto ensure_trailing_slash = []( std :: string p ) {
// Make sure to add trailing slash
if ( p . back () != DIRECTORY_SEPARATOR ) {
p += DIRECTORY_SEPARATOR ;
}
return p ;
};
2024-05-22 20:04:20 +03:00
if ( getenv ( "LLAMA_CACHE" )) {
cache_directory = std :: getenv ( "LLAMA_CACHE" );
} else {
2026-02-14 09:47:01 +01:00
#if defined(__linux__) || defined(__FreeBSD__) || defined(_AIX) || \
defined(__OpenBSD__) || defined(__NetBSD__)
2024-05-22 20:04:20 +03:00
if ( std :: getenv ( "XDG_CACHE_HOME" )) {
cache_directory = std :: getenv ( "XDG_CACHE_HOME" );
2025-09-27 02:03:33 +08:00
} else if ( std :: getenv ( "HOME" )) {
2024-05-22 20:04:20 +03:00
cache_directory = std :: getenv ( "HOME" ) + std :: string ( "/.cache/" );
2025-09-27 02:03:33 +08:00
} else {
#if defined(__linux__)
/* no $HOME is defined, fallback to getpwuid */
struct passwd * pw = getpwuid ( getuid ());
if (( ! pw ) || ( ! pw -> pw_dir )) {
throw std :: runtime_error ( "Failed to find $HOME directory" );
}
cache_directory = std :: string ( pw -> pw_dir ) + std :: string ( "/.cache/" );
#else /* defined(__linux__) */
throw std :: runtime_error ( "Failed to find $HOME directory" );
#endif /* defined(__linux__) */
2024-05-22 20:04:20 +03:00
}
#elif defined(__APPLE__)
cache_directory = std :: getenv ( "HOME" ) + std :: string ( "/Library/Caches/" );
#elif defined(_WIN32)
2024-05-25 05:30:59 +02:00
cache_directory = std :: getenv ( "LOCALAPPDATA" );
2025-12-03 01:25:34 -08:00
#elif defined(__EMSCRIPTEN__)
GGML_ABORT ( "not implemented on this platform" );
2025-04-11 12:45:44 -07:00
#else
# error Unknown architecture
#endif
2024-05-25 05:30:59 +02:00
cache_directory = ensure_trailing_slash ( cache_directory );
2024-05-22 20:04:20 +03:00
cache_directory += "llama.cpp" ;
2023-12-05 15:05:51 +05:00
}
2024-05-25 05:30:59 +02:00
return ensure_trailing_slash ( cache_directory );
2023-12-05 15:05:51 +05:00
}
2024-06-08 20:21:08 +01:00
std :: string fs_get_cache_file ( const std :: string & filename ) {
GGML_ASSERT ( filename . find ( DIRECTORY_SEPARATOR ) == std :: string :: npos );
std :: string cache_directory = fs_get_cache_directory ();
const bool success = fs_create_directory_with_parents ( cache_directory );
if ( ! success ) {
throw std :: runtime_error ( "failed to create cache directory: " + cache_directory );
}
return cache_directory + filename ;
}
2025-12-01 19:41:04 +01:00
std :: vector < common_file_info > fs_list ( const std :: string & path , bool include_directories ) {
2025-11-08 21:54:14 +01:00
std :: vector < common_file_info > files ;
if ( path . empty ()) return files ;
std :: filesystem :: path dir ( path );
if ( ! std :: filesystem :: exists ( dir ) || ! std :: filesystem :: is_directory ( dir )) {
return files ;
}
for ( const auto & entry : std :: filesystem :: directory_iterator ( dir )) {
try {
// Only include regular files (skip directories)
const auto & p = entry . path ();
if ( std :: filesystem :: is_regular_file ( p )) {
common_file_info info ;
2025-12-01 19:41:04 +01:00
info . path = p . string ();
info . name = p . filename (). string ();
info . is_dir = false ;
2025-11-08 21:54:14 +01:00
try {
info . size = static_cast < size_t > ( std :: filesystem :: file_size ( p ));
} catch ( const std :: filesystem :: filesystem_error & ) {
info . size = 0 ;
}
files . push_back ( std :: move ( info ));
2025-12-01 19:41:04 +01:00
} else if ( include_directories && std :: filesystem :: is_directory ( p )) {
common_file_info info ;
info . path = p . string ();
info . name = p . filename (). string ();
info . size = 0 ; // Directories have no size
info . is_dir = true ;
files . push_back ( std :: move ( info ));
2025-11-08 21:54:14 +01:00
}
} catch ( const std :: filesystem :: filesystem_error & ) {
// skip entries we cannot inspect
continue ;
}
}
return files ;
}
2026-06-19 22:28:38 +02:00
std :: ifstream fs_open_ifstream ( const std :: string & fname , std :: ios_base :: openmode mode ) {
#ifdef _WIN32
int wlen = MultiByteToWideChar ( CP_UTF8 , 0 , fname . c_str (), - 1 , NULL , 0 );
if ( ! wlen ) { return std :: ifstream (); }
std :: vector < wchar_t > wfname ( wlen );
( void ) MultiByteToWideChar ( CP_UTF8 , 0 , fname . c_str (), - 1 , wfname . data (), wlen );
return std :: ifstream ( wfname . data (), mode );
#else
return std :: ifstream ( fname , mode );
#endif
}
2025-12-07 03:43:50 +01:00
//
// TTY utils
//
bool tty_can_use_colors () {
// Check NO_COLOR environment variable (https://no-color.org/)
if ( const char * no_color = std :: getenv ( "NO_COLOR" )) {
if ( no_color [ 0 ] != '\0' ) {
return false ;
}
}
// Check TERM environment variable
if ( const char * term = std :: getenv ( "TERM" )) {
if ( std :: strcmp ( term , "dumb" ) == 0 ) {
return false ;
}
}
// Check if stdout and stderr are connected to a terminal
// We check both because log messages can go to either
bool stdout_is_tty = isatty ( fileno ( stdout ));
bool stderr_is_tty = isatty ( fileno ( stderr ));
return stdout_is_tty || stderr_is_tty ;
}
2024-05-22 20:04:20 +03:00
2023-08-21 23:07:43 +03:00
//
// Model utils
//
2025-04-01 23:44:05 +02:00
2025-12-14 10:11:13 +02:00
// TODO: move to common/sampling
static void common_init_sampler_from_model (
2025-11-25 09:56:07 +08:00
const llama_model * model ,
common_params_sampling & sparams ) {
const uint64_t config = sparams . user_sampling_config ;
auto get_int32 = [ & ]( const char * key , int32_t & dst , uint64_t user_config ) {
2025-12-14 10:11:13 +02:00
if ( config & user_config ) {
return ;
}
2025-11-25 09:56:07 +08:00
char buf [ 64 ] = { 0 };
if ( llama_model_meta_val_str ( model , key , buf , sizeof ( buf )) > 0 ) {
char * end = nullptr ;
int32_t v = strtol ( buf , & end , 10 );
2025-12-14 10:11:13 +02:00
if ( end && end != buf ) {
dst = v ;
}
2025-11-25 09:56:07 +08:00
}
};
auto get_float = [ & ]( const char * key , float & dst , uint64_t user_config ) {
2025-12-14 10:11:13 +02:00
if ( config & user_config ) {
return ;
}
2025-11-25 09:56:07 +08:00
char buf [ 128 ] = { 0 };
if ( llama_model_meta_val_str ( model , key , buf , sizeof ( buf )) > 0 ) {
char * end = nullptr ;
float v = strtof ( buf , & end );
2025-12-14 10:11:13 +02:00
if ( end && end != buf ) {
dst = v ;
}
2025-11-25 09:56:07 +08:00
}
};
// Sampling sequence
if ( ! ( config & common_params_sampling_config :: COMMON_PARAMS_SAMPLING_CONFIG_SAMPLERS )) {
char buf [ 512 ] = { 0 };
if ( llama_model_meta_val_str ( model , llama_model_meta_key_str ( LLAMA_MODEL_META_KEY_SAMPLING_SEQUENCE ), buf , sizeof ( buf )) > 0 ) {
const std :: vector < std :: string > sampler_names = string_split < std :: string > ( std :: string ( buf ), ';' );
if ( ! sampler_names . empty ()) {
2026-06-07 15:48:11 -05:00
sparams . samplers = common_sampler_types_from_names ( sampler_names );
2025-11-25 09:56:07 +08:00
}
}
}
get_int32 ( llama_model_meta_key_str ( LLAMA_MODEL_META_KEY_SAMPLING_TOP_K ), sparams . top_k , common_params_sampling_config :: COMMON_PARAMS_SAMPLING_CONFIG_TOP_K );
get_float ( llama_model_meta_key_str ( LLAMA_MODEL_META_KEY_SAMPLING_TOP_P ), sparams . top_p , common_params_sampling_config :: COMMON_PARAMS_SAMPLING_CONFIG_TOP_P );
get_float ( llama_model_meta_key_str ( LLAMA_MODEL_META_KEY_SAMPLING_MIN_P ), sparams . min_p , common_params_sampling_config :: COMMON_PARAMS_SAMPLING_CONFIG_MIN_P );
get_float ( llama_model_meta_key_str ( LLAMA_MODEL_META_KEY_SAMPLING_XTC_PROBABILITY ), sparams . xtc_probability , common_params_sampling_config :: COMMON_PARAMS_SAMPLING_CONFIG_XTC_PROBABILITY );
get_float ( llama_model_meta_key_str ( LLAMA_MODEL_META_KEY_SAMPLING_XTC_THRESHOLD ), sparams . xtc_threshold , common_params_sampling_config :: COMMON_PARAMS_SAMPLING_CONFIG_XTC_THRESHOLD );
get_float ( llama_model_meta_key_str ( LLAMA_MODEL_META_KEY_SAMPLING_TEMP ), sparams . temp , common_params_sampling_config :: COMMON_PARAMS_SAMPLING_CONFIG_TEMP );
get_int32 ( llama_model_meta_key_str ( LLAMA_MODEL_META_KEY_SAMPLING_PENALTY_LAST_N ), sparams . penalty_last_n , common_params_sampling_config :: COMMON_PARAMS_SAMPLING_CONFIG_PENALTY_LAST_N );
get_float ( llama_model_meta_key_str ( LLAMA_MODEL_META_KEY_SAMPLING_PENALTY_REPEAT ), sparams . penalty_repeat , common_params_sampling_config :: COMMON_PARAMS_SAMPLING_CONFIG_PENALTY_REPEAT );
get_int32 ( llama_model_meta_key_str ( LLAMA_MODEL_META_KEY_SAMPLING_MIROSTAT ), sparams . mirostat , common_params_sampling_config :: COMMON_PARAMS_SAMPLING_CONFIG_MIROSTAT );
get_float ( llama_model_meta_key_str ( LLAMA_MODEL_META_KEY_SAMPLING_MIROSTAT_TAU ), sparams . mirostat_tau , common_params_sampling_config :: COMMON_PARAMS_SAMPLING_CONFIG_MIROSTAT_TAU );
get_float ( llama_model_meta_key_str ( LLAMA_MODEL_META_KEY_SAMPLING_MIROSTAT_ETA ), sparams . mirostat_eta , common_params_sampling_config :: COMMON_PARAMS_SAMPLING_CONFIG_MIROSTAT_ETA );
}
2025-12-14 10:11:13 +02:00
struct common_init_result :: impl {
impl () = default ;
~ impl () = default ;
2025-12-22 11:00:37 +01:00
// note: the order in which model, context, etc. are declared matters because their destructors will be called bottom-to-top
2025-12-14 10:11:13 +02:00
llama_model_ptr model ;
llama_context_ptr context ;
std :: vector < llama_adapter_lora_ptr > lora ;
std :: vector < common_sampler_ptr > samplers ;
2026-01-04 21:22:16 +01:00
std :: vector < llama_sampler_seq_config > samplers_seq_config ;
2025-12-14 10:11:13 +02:00
};
2026-05-19 09:46:34 +03:00
common_init_result :: common_init_result ( common_params & params , bool model_only ) :
2025-12-14 10:11:13 +02:00
pimpl ( new impl {}) {
2025-12-15 09:24:59 +01:00
auto mparams = common_model_params_to_llama ( params );
auto cparams = common_context_params_to_llama ( params );
if ( params . fit_params ) {
2026-06-28 08:52:15 +03:00
COM_TRC ( "%s" , "fitting params to device memory ... \n " );
COM_TRC ( "%s" , "(for bugs during this step try to reproduce them with -fit off, or provide --verbose logs if the bug only occurs with -fit on) \n " );
2026-04-21 09:54:36 +03:00
common_fit_params ( params . model . path . c_str (), & mparams , & cparams ,
2026-01-28 18:42:42 +01:00
params . tensor_split ,
params . tensor_buft_overrides . data (),
params . fit_params_target . data (),
params . fit_params_min_ctx ,
2026-05-19 21:33:23 +02:00
params . verbosity >= LOG_LEVEL_DEBUG ? GGML_LOG_LEVEL_DEBUG : GGML_LOG_LEVEL_ERROR );
2025-12-15 09:24:59 +01:00
}
2024-05-22 20:04:20 +03:00
2025-04-01 23:44:05 +02:00
llama_model * model = llama_model_load_from_file ( params . model . path . c_str (), mparams );
2024-05-22 20:04:20 +03:00
if ( model == NULL ) {
2025-12-14 10:11:13 +02:00
return ;
2024-05-22 20:04:20 +03:00
}
2025-12-14 10:11:13 +02:00
pimpl -> model . reset ( model );
2025-11-25 09:56:07 +08:00
2026-05-19 09:46:34 +03:00
if ( model_only ) {
return ;
}
2025-01-12 11:32:42 +02:00
const llama_vocab * vocab = llama_model_get_vocab ( model );
2026-03-18 11:03:26 +01:00
// load and optionally apply lora adapters
2025-12-30 15:53:12 +01:00
for ( auto & la : params . lora_adapters ) {
llama_adapter_lora_ptr lora ;
lora . reset ( llama_adapter_lora_init ( model , la . path . c_str ()));
if ( lora == nullptr ) {
2026-06-28 08:52:15 +03:00
COM_ERR ( "failed to load lora adapter '%s' \n " , la . path . c_str ());
2025-12-30 15:53:12 +01:00
return ;
}
char buf [ 1024 ];
la . ptr = lora . get ();
llama_adapter_meta_val_str ( la . ptr , "adapter.lora.task_name" , buf , sizeof ( buf ));
la . task_name = buf ;
llama_adapter_meta_val_str ( la . ptr , "adapter.lora.prompt_prefix" , buf , sizeof ( buf ));
la . prompt_prefix = buf ;
pimpl -> lora . emplace_back ( std :: move ( lora )); // copy to list of loaded adapters
}
2025-12-14 10:11:13 +02:00
// updates params.sampling
// TODO: fix naming
common_init_sampler_from_model ( model , params . sampling );
if ( params . sampling . ignore_eos && llama_vocab_eos ( vocab ) == LLAMA_TOKEN_NULL ) {
2026-06-28 08:52:15 +03:00
COM_WRN ( "%s" , "vocab does not have an EOS token, ignoring --ignore-eos \n " );
2025-12-14 10:11:13 +02:00
params . sampling . ignore_eos = false ;
}
// initialize once
for ( llama_token i = 0 ; i < llama_vocab_n_tokens ( vocab ); i ++ ) {
if ( llama_vocab_is_eog ( vocab , i )) {
2026-06-28 08:52:15 +03:00
COM_TRC ( "added %s logit bias = %f \n " , common_token_to_piece ( vocab , i ). c_str (), - INFINITY );
2025-12-14 10:11:13 +02:00
params . sampling . logit_bias_eog . push_back ({ i , - INFINITY });
}
}
if ( params . sampling . ignore_eos ) {
// add EOG biases to the active set of logit biases
params . sampling . logit_bias . insert (
params . sampling . logit_bias . end (),
params . sampling . logit_bias_eog . begin (), params . sampling . logit_bias_eog . end ());
}
2026-01-04 21:22:16 +01:00
// init the backend samplers as part of the context creation
2025-12-14 10:11:13 +02:00
pimpl -> samplers . resize ( cparams . n_seq_max );
2026-01-04 21:22:16 +01:00
pimpl -> samplers_seq_config . resize ( cparams . n_seq_max );
2025-12-14 10:11:13 +02:00
for ( int i = 0 ; i < ( int ) cparams . n_seq_max ; ++ i ) {
2026-08-04 20:28:55 +02:00
pimpl -> samplers [ i ]. reset ( common_sampler_init ( model , params . sampling ));
2026-01-04 21:22:16 +01:00
pimpl -> samplers_seq_config [ i ] = { i , common_sampler_get ( pimpl -> samplers [ i ]. get ()) };
}
if ( params . sampling . backend_sampling ) {
cparams . samplers = pimpl -> samplers_seq_config . data ();
cparams . n_samplers = pimpl -> samplers_seq_config . size ();
2025-12-14 10:11:13 +02:00
}
2025-01-12 11:32:42 +02:00
llama_context * lctx = llama_init_from_model ( model , cparams );
2024-05-22 20:04:20 +03:00
if ( lctx == NULL ) {
2026-06-28 08:52:15 +03:00
COM_ERR ( "failed to create context with model '%s' \n " , params . model . path . c_str ());
2025-12-14 10:11:13 +02:00
return ;
2024-05-22 20:04:20 +03:00
}
2025-12-14 10:11:13 +02:00
pimpl -> context . reset ( lctx );
}
llama_model * common_init_result :: model () {
return pimpl -> model . get ();
}
llama_context * common_init_result :: context () {
return pimpl -> context . get ();
}
common_sampler * common_init_result :: sampler ( llama_seq_id seq_id ) {
2026-03-31 07:04:42 -03:00
if ( seq_id < 0 || seq_id >= ( int ) pimpl -> samplers . size ()) {
return nullptr ;
}
2025-12-14 10:11:13 +02:00
return pimpl -> samplers [ seq_id ]. get ();
}
2026-01-04 21:22:16 +01:00
void common_init_result :: reset_samplers () {
for ( int i = 0 ; i < ( int ) pimpl -> samplers . size (); ++ i ) {
llama_sampler_reset ( common_sampler_get ( pimpl -> samplers [ i ]. get ()));
}
}
2025-12-14 10:11:13 +02:00
std :: vector < llama_adapter_lora_ptr > & common_init_result :: lora () {
return pimpl -> lora ;
}
2026-05-19 09:46:34 +03:00
common_init_result_ptr common_init_from_params ( common_params & params , bool model_only ) {
common_init_result_ptr res ( new common_init_result ( params , model_only ));
2025-12-14 10:11:13 +02:00
llama_model * model = res -> model ();
if ( model == NULL ) {
2026-06-28 08:52:15 +03:00
COM_ERR ( "failed to load model '%s' \n " , params . model . path . c_str ());
2025-12-14 10:11:13 +02:00
return res ;
}
2026-05-19 09:46:34 +03:00
if ( model_only ) {
return res ;
}
2025-12-14 10:11:13 +02:00
llama_context * lctx = res -> context ();
if ( lctx == NULL ) {
2026-06-28 08:52:15 +03:00
COM_ERR ( "failed to create context with model '%s' \n " , params . model . path . c_str ());
2025-12-14 10:11:13 +02:00
return res ;
}
const llama_vocab * vocab = llama_model_get_vocab ( model );
2025-06-06 14:11:15 +03:00
if ( params . ctx_shift && ! llama_memory_can_shift ( llama_get_memory ( lctx ))) {
2026-06-28 08:52:15 +03:00
COM_WRN ( "%s" , "KV cache shifting is not supported for this context, disabling KV cache shifting \n " );
2025-01-03 20:13:18 +08:00
params . ctx_shift = false ;
2024-11-19 13:29:26 +02:00
}
2024-05-22 20:04:20 +03:00
if ( ! params . control_vectors . empty ()) {
if ( params . control_vector_layer_start <= 0 ) params . control_vector_layer_start = 1 ;
2025-01-12 11:32:42 +02:00
if ( params . control_vector_layer_end <= 0 ) params . control_vector_layer_end = llama_model_n_layer ( model );
2024-05-22 20:04:20 +03:00
2024-10-10 22:57:42 +02:00
const auto cvec = common_control_vector_load ( params . control_vectors );
2024-05-22 20:04:20 +03:00
if ( cvec . n_embd == - 1 ) {
2025-12-14 10:11:13 +02:00
return res ;
2024-05-22 20:04:20 +03:00
}
2026-02-14 03:06:27 -05:00
int err = llama_set_adapter_cvec (
2025-01-12 11:32:42 +02:00
lctx ,
cvec . data . data (),
cvec . data . size (),
cvec . n_embd ,
params . control_vector_layer_start ,
params . control_vector_layer_end );
2024-05-22 20:04:20 +03:00
if ( err ) {
2025-12-14 10:11:13 +02:00
return res ;
2024-05-22 20:04:20 +03:00
}
}
2025-06-16 14:14:00 +03:00
if ( llama_pooling_type ( lctx ) == LLAMA_POOLING_TYPE_RANK ) {
bool ok = true ;
if ( llama_vocab_bos ( vocab ) == LLAMA_TOKEN_NULL ) {
2026-06-28 08:52:15 +03:00
COM_WRN ( "%s" , "vocab does not have a BOS token, reranking will not work \n " );
2025-06-16 14:14:00 +03:00
ok = false ;
}
bool has_eos = llama_vocab_eos ( vocab ) != LLAMA_TOKEN_NULL ;
bool has_sep = llama_vocab_sep ( vocab ) != LLAMA_TOKEN_NULL ;
2025-09-25 03:53:09 -05:00
bool has_rerank_prompt = llama_model_chat_template ( model , "rerank" ) != NULL ;
2025-06-16 14:14:00 +03:00
2025-09-25 03:53:09 -05:00
if ( ! has_eos && ! has_sep && ! has_rerank_prompt ) {
2026-06-28 08:52:15 +03:00
COM_WRN ( "%s" , "vocab does not have an EOS token, SEP token, or rerank prompt. Reranking will not work \n " );
2025-06-16 14:14:00 +03:00
ok = false ;
} else if ( ! has_eos ) {
2026-06-28 08:52:15 +03:00
COM_WRN ( "%s" , "vocab does not have an EOS token, using SEP token as fallback \n " );
2025-06-16 14:14:00 +03:00
}
if ( ! ok ) {
2025-12-14 10:11:13 +02:00
return res ;
2025-06-16 14:14:00 +03:00
}
}
2024-08-06 17:33:39 +02:00
if ( ! params . lora_init_without_apply ) {
2025-01-12 11:32:42 +02:00
common_set_adapter_lora ( lctx , params . lora_adapters );
2024-05-22 20:04:20 +03:00
}
if ( params . warmup ) {
2026-06-28 08:52:15 +03:00
COM_TRC ( "%s" , "warming up the model with an empty run - please wait ... (--no-warmup to disable) \n " );
2024-05-22 20:04:20 +03:00
2024-07-04 15:46:11 +02:00
std :: vector < llama_token > tmp ;
2025-01-12 11:32:42 +02:00
llama_token bos = llama_vocab_bos ( vocab );
llama_token eos = llama_vocab_eos ( vocab );
2024-07-04 15:46:11 +02:00
// some models (e.g. T5) don't have a BOS token
2024-09-08 00:33:13 +03:00
if ( bos != LLAMA_TOKEN_NULL ) {
2024-07-04 15:46:11 +02:00
tmp . push_back ( bos );
}
2024-09-08 00:33:13 +03:00
if ( eos != LLAMA_TOKEN_NULL ) {
tmp . push_back ( eos );
}
if ( tmp . empty ()) {
tmp . push_back ( 0 );
}
2024-07-04 15:46:11 +02:00
if ( llama_model_has_encoder ( model )) {
2024-10-18 23:18:01 +02:00
llama_encode ( lctx , llama_batch_get_one ( tmp . data (), tmp . size ()));
2024-07-04 15:46:11 +02:00
llama_token decoder_start_token_id = llama_model_decoder_start_token ( model );
2025-01-06 10:52:15 +02:00
if ( decoder_start_token_id == LLAMA_TOKEN_NULL ) {
2024-07-04 15:46:11 +02:00
decoder_start_token_id = bos ;
}
tmp . clear ();
tmp . push_back ( decoder_start_token_id );
}
2024-08-10 11:43:26 +02:00
if ( llama_model_has_decoder ( model )) {
2024-10-18 23:18:01 +02:00
llama_decode ( lctx , llama_batch_get_one ( tmp . data (), std :: min ( tmp . size (), ( size_t ) params . n_batch )));
2024-08-10 11:43:26 +02:00
}
2025-06-06 14:11:15 +03:00
llama_memory_clear ( llama_get_memory ( lctx ), true );
2024-05-22 20:04:20 +03:00
llama_synchronize ( lctx );
2024-09-13 09:53:38 +03:00
llama_perf_context_reset ( lctx );
2026-01-04 21:22:16 +01:00
// reset samplers to reset RNG state after warmup to the seeded state
res -> reset_samplers ();
2024-05-22 20:04:20 +03:00
}
2025-12-14 10:11:13 +02:00
return res ;
2024-05-22 20:04:20 +03:00
}
2025-12-14 10:11:13 +02:00
common_init_result ::~ common_init_result () = default ;
2026-04-20 08:42:37 +03:00
std :: string common_get_model_endpoint () {
2026-08-03 18:58:15 +02:00
std :: string endpoint = common_get_env ( "MODEL_ENDPOINT" );
if ( endpoint . empty ()) {
// the HF_ENDPOINT variable is respected for backward compatibility
endpoint = common_get_env ( "HF_ENDPOINT" );
2025-04-11 20:01:56 +08:00
}
2026-08-03 18:58:15 +02:00
if ( endpoint . empty ()) {
return "https://huggingface.co/" ;
}
if ( endpoint . back () != '/' ) {
endpoint += '/' ;
}
return endpoint ;
2025-04-11 20:01:56 +08:00
}
2026-07-30 19:34:04 +03:00
char * common_get_model_or_exit ( int argc , char * argv []) {
if ( argc > 1 ) {
return argv [ 1 ];
}
char * path = getenv ( "LLAMACPP_TEST_MODELFILE" );
if ( ! path || strlen ( path ) == 0 ) {
fprintf ( stderr , " \033 [33mWARNING: No model file provided. Skipping this test. Set LLAMACPP_TEST_MODELFILE=<gguf_model_path> to silence this warning and run this test. \n\033 [0m" );
exit ( EXIT_SUCCESS );
}
return path ;
}
2026-04-20 08:42:37 +03:00
common_context_seq_rm_type common_context_can_seq_rm ( llama_context * ctx ) {
auto * mem = llama_get_memory ( ctx );
if ( mem == nullptr ) {
return COMMON_CONTEXT_SEQ_RM_TYPE_NO ;
}
common_context_seq_rm_type res = COMMON_CONTEXT_SEQ_RM_TYPE_PART ;
llama_memory_clear ( mem , true );
// eval 2 tokens to check if the context is compatible
std :: vector < llama_token > tmp ;
tmp . push_back ( 0 );
tmp . push_back ( 0 );
int ret = llama_decode ( ctx , llama_batch_get_one ( tmp . data (), tmp . size ()));
if ( ret != 0 ) {
2026-06-28 08:52:15 +03:00
COM_ERR ( "llama_decode() failed: %d \n " , ret );
2026-04-20 08:42:37 +03:00
res = COMMON_CONTEXT_SEQ_RM_TYPE_NO ;
goto done ;
}
2026-05-16 20:06:23 +08:00
if ( llama_n_rs_seq ( ctx ) > 0 ) {
2026-06-28 08:52:15 +03:00
COM_TRC ( "%s" , "the context supports bounded partial sequence removal \n " );
2026-05-16 20:06:23 +08:00
res = COMMON_CONTEXT_SEQ_RM_TYPE_RS ;
goto done ;
}
2026-04-20 08:42:37 +03:00
// try to remove the last tokens
if ( ! llama_memory_seq_rm ( mem , 0 , 1 , - 1 )) {
2026-06-28 08:52:15 +03:00
COM_TRC ( "%s" , "the context does not support partial sequence removal \n " );
2026-04-20 08:42:37 +03:00
res = COMMON_CONTEXT_SEQ_RM_TYPE_FULL ;
goto done ;
}
done :
llama_memory_clear ( mem , true );
llama_synchronize ( ctx );
return res ;
}
2026-07-28 16:35:20 +02:00
static void common_context_seq_rm ( llama_context * ctx , llama_seq_id seq_id , llama_pos p0 , llama_pos p1 ) {
2026-05-16 20:06:23 +08:00
auto * mem = llama_get_memory ( ctx );
if ( ! llama_memory_seq_rm ( mem , seq_id , p0 , p1 )) {
GGML_ABORT ( "%s" , string_format ( "failed to remove sequence %d with p0=%d, p1=%d \n " , seq_id , p0 , p1 ). c_str ());
}
}
2026-07-28 16:35:20 +02:00
static void common_context_seq_cp ( llama_context * ctx , llama_seq_id seq_id_src , llama_seq_id seq_id_dst , llama_pos p0 , llama_pos p1 ) {
2026-05-16 20:06:23 +08:00
auto * mem = llama_get_memory ( ctx );
llama_memory_seq_cp ( mem , seq_id_src , seq_id_dst , p0 , p1 );
}
2026-07-28 16:35:20 +02:00
static void common_context_seq_add ( llama_context * ctx , llama_seq_id seq_id , llama_pos p0 , llama_pos p1 , llama_pos delta ) {
2026-05-16 20:06:23 +08:00
auto * mem = llama_get_memory ( ctx );
llama_memory_seq_add ( mem , seq_id , p0 , p1 , delta );
}
2026-07-28 16:35:20 +02:00
void common_memory :: init ( llama_context * ctx_tgt , llama_context * ctx_dft ) {
this -> ctx_tgt = ctx_tgt ;
this -> ctx_dft = ctx_dft ;
}
void common_memory :: seq_rm ( llama_seq_id seq_id , llama_pos p0 , llama_pos p1 ) const {
common_context_seq_rm ( ctx_tgt , seq_id , p0 , p1 );
if ( ctx_dft ) {
common_context_seq_rm ( ctx_dft , seq_id , p0 , p1 );
}
}
void common_memory :: seq_cp ( llama_seq_id seq_id_src , llama_seq_id seq_id_dst , llama_pos p0 , llama_pos p1 ) const {
common_context_seq_cp ( ctx_tgt , seq_id_src , seq_id_dst , p0 , p1 );
if ( ctx_dft ) {
common_context_seq_cp ( ctx_dft , seq_id_src , seq_id_dst , p0 , p1 );
}
}
void common_memory :: seq_add ( llama_seq_id seq_id , llama_pos p0 , llama_pos p1 , llama_pos delta ) const {
common_context_seq_add ( ctx_tgt , seq_id , p0 , p1 , delta );
if ( ctx_dft ) {
common_context_seq_add ( ctx_dft , seq_id , p0 , p1 , delta );
}
}
2025-01-12 11:32:42 +02:00
void common_set_adapter_lora ( struct llama_context * ctx , std :: vector < common_adapter_lora_info > & lora ) {
2026-02-14 03:06:27 -05:00
std :: vector < llama_adapter_lora *> loras ;
std :: vector < float > scales ;
for ( auto & la : lora ) {
loras . push_back ( la . ptr );
scales . push_back ( la . scale );
2024-08-06 17:33:39 +02:00
}
2026-02-14 03:06:27 -05:00
llama_set_adapters_lora ( ctx , loras . data (), loras . size (), scales . data ());
2024-08-06 17:33:39 +02:00
}
2024-11-25 19:30:06 +01:00
struct llama_model_params common_model_params_to_llama ( common_params & params ) {
2023-09-28 21:42:38 +02:00
auto mparams = llama_model_default_params ();
2023-05-02 22:39:51 +02:00
2024-11-25 19:30:06 +01:00
if ( ! params . devices . empty ()) {
mparams . devices = params . devices . data ();
}
2025-04-02 14:52:01 +02:00
2025-12-27 20:18:35 +01:00
mparams . n_gpu_layers = params . n_gpu_layers ;
2023-09-28 21:42:38 +02:00
mparams . main_gpu = params . main_gpu ;
2024-01-12 20:07:38 +01:00
mparams . split_mode = params . split_mode ;
2026-07-23 20:32:56 +08:00
mparams . load_mode = params . load_mode ;
2023-09-28 21:42:38 +02:00
mparams . tensor_split = params . tensor_split ;
2024-04-26 18:39:58 +02:00
mparams . check_tensors = params . check_tensors ;
2025-07-31 09:11:34 -07:00
mparams . use_extra_bufts = ! params . no_extra_bufts ;
2025-10-06 12:55:53 -05:00
mparams . no_host = params . no_host ;
2025-04-02 14:52:01 +02:00
2023-12-05 10:19:18 -07:00
if ( params . kv_overrides . empty ()) {
mparams . kv_overrides = NULL ;
} else {
GGML_ASSERT ( params . kv_overrides . back (). key [ 0 ] == 0 && "KV overrides not terminated with empty key" );
mparams . kv_overrides = params . kv_overrides . data ();
}
2023-05-02 22:39:51 +02:00
2025-04-02 14:52:01 +02:00
if ( params . tensor_buft_overrides . empty ()) {
mparams . tensor_buft_overrides = NULL ;
} else {
GGML_ASSERT ( params . tensor_buft_overrides . back (). pattern == nullptr && "Tensor buffer overrides not terminated with empty pattern" );
mparams . tensor_buft_overrides = params . tensor_buft_overrides . data ();
}
2025-05-19 21:17:36 +02:00
mparams . progress_callback = params . load_progress_callback ;
mparams . progress_callback_user_data = params . load_progress_callback_user_data ;
2026-04-02 18:19:20 +02:00
mparams . no_alloc = params . no_alloc ;
2026-07-31 14:57:02 +02:00
mparams . load_mtp = std :: find ( params . speculative . types . begin (), params . speculative . types . end (), COMMON_SPECULATIVE_TYPE_DRAFT_MTP ) != params . speculative . types . end ();
2025-05-19 21:17:36 +02:00
2023-09-28 21:42:38 +02:00
return mparams ;
}
2024-10-10 22:57:42 +02:00
struct llama_context_params common_context_params_to_llama ( const common_params & params ) {
2023-09-28 21:42:38 +02:00
auto cparams = llama_context_default_params ();
2023-11-01 18:04:33 -04:00
cparams . n_ctx = params . n_ctx ;
2024-03-11 17:49:47 +02:00
cparams . n_seq_max = params . n_parallel ;
2026-05-16 20:06:23 +08:00
cparams . n_rs_seq = params . speculative . need_n_rs_seq ();
2026-06-01 23:01:38 +08:00
cparams . n_outputs_max = std :: max ( params . n_outputs_max , 0 );
2026-08-10 19:28:56 +05:30
cparams . n_outputs_max_per_seq = std :: max ( params . n_outputs_max_per_seq , 0 );
2024-03-13 18:54:21 +01:00
cparams . n_batch = params . n_batch ;
cparams . n_ubatch = params . n_ubatch ;
2024-08-29 19:20:53 -04:00
cparams . n_threads = params . cpuparams . n_threads ;
cparams . n_threads_batch = params . cpuparams_batch . n_threads == - 1 ?
2024-10-21 16:20:46 +03:00
params . cpuparams . n_threads : params . cpuparams_batch . n_threads ;
2024-03-04 22:31:20 +02:00
cparams . embeddings = params . embedding ;
2023-11-01 18:04:33 -04:00
cparams . rope_scaling_type = params . rope_scaling_type ;
cparams . rope_freq_base = params . rope_freq_base ;
cparams . rope_freq_scale = params . rope_freq_scale ;
cparams . yarn_ext_factor = params . yarn_ext_factor ;
cparams . yarn_attn_factor = params . yarn_attn_factor ;
cparams . yarn_beta_fast = params . yarn_beta_fast ;
cparams . yarn_beta_slow = params . yarn_beta_slow ;
cparams . yarn_orig_ctx = params . yarn_orig_ctx ;
2024-03-03 04:40:27 -06:00
cparams . pooling_type = params . pooling_type ;
2024-07-05 02:05:56 -05:00
cparams . attention_type = params . attention_type ;
2025-08-30 16:32:10 +02:00
cparams . flash_attn_type = params . flash_attn_type ;
2024-04-11 14:51:07 +02:00
cparams . cb_eval = params . cb_eval ;
cparams . cb_eval_user_data = params . cb_eval_user_data ;
2023-12-07 13:03:17 +02:00
cparams . offload_kqv = ! params . no_kv_offload ;
2024-09-13 09:53:38 +03:00
cparams . no_perf = params . no_perf ;
2025-05-11 20:18:39 +08:00
cparams . op_offload = ! params . no_op_offload ;
2025-05-20 08:05:46 +03:00
cparams . swa_full = params . swa_full ;
2025-07-16 16:35:42 +03:00
cparams . kv_unified = params . kv_unified ;
2023-12-07 13:03:17 +02:00
2024-12-12 22:53:05 +01:00
cparams . type_k = params . cache_type_k ;
cparams . type_v = params . cache_type_v ;
2023-09-28 21:42:38 +02:00
return cparams ;
2023-07-12 00:18:43 +08:00
}
2026-04-28 09:07:33 +03:00
struct ggml_threadpool_params ggml_threadpool_params_from_cpu_params ( const common_cpu_params & params ) {
2024-08-29 19:20:53 -04:00
struct ggml_threadpool_params tpp ;
ggml_threadpool_params_init ( & tpp , params . n_threads ); // setup the defaults
if ( params . mask_valid ) {
std :: memcpy ( & tpp . cpumask , & params . cpumask , GGML_MAX_N_THREADS );
}
tpp . prio = params . priority ;
tpp . poll = params . poll ;
tpp . strict_cpu = params . strict_cpu ;
return tpp ;
}
2024-05-22 20:04:20 +03:00
//
// Batch utils
//
2023-07-12 00:18:43 +08:00
2024-10-10 22:57:42 +02:00
void common_batch_clear ( struct llama_batch & batch ) {
2024-05-22 20:04:20 +03:00
batch . n_tokens = 0 ;
}
2024-03-22 15:33:38 +02:00
2024-10-10 22:57:42 +02:00
void common_batch_add (
2024-05-22 20:04:20 +03:00
struct llama_batch & batch ,
llama_token id ,
llama_pos pos ,
const std :: vector < llama_seq_id > & seq_ids ,
bool logits ) {
2024-09-29 05:25:00 -07:00
GGML_ASSERT ( batch . seq_id [ batch . n_tokens ] && "llama_batch size exceeded" );
2024-05-22 20:04:20 +03:00
batch . token [ batch . n_tokens ] = id ;
batch . pos [ batch . n_tokens ] = pos ;
batch . n_seq_id [ batch . n_tokens ] = seq_ids . size ();
for ( size_t i = 0 ; i < seq_ids . size (); ++ i ) {
batch . seq_id [ batch . n_tokens ][ i ] = seq_ids [ i ];
2024-03-17 19:12:37 +01:00
}
2024-05-22 20:04:20 +03:00
batch . logits [ batch . n_tokens ] = logits ;
2024-03-22 15:33:38 +02:00
2024-05-22 20:04:20 +03:00
batch . n_tokens ++ ;
2023-05-02 22:39:51 +02:00
}
2023-08-21 23:07:43 +03:00
//
// Vocab utils
//
2024-10-10 22:57:42 +02:00
std :: vector < llama_token > common_tokenize (
2023-09-28 21:42:38 +02:00
const struct llama_context * ctx ,
const std :: string & text ,
2024-04-09 13:44:08 -04:00
bool add_special ,
bool parse_special ) {
2025-01-12 11:32:42 +02:00
const llama_model * model = llama_get_model ( ctx );
const llama_vocab * vocab = llama_model_get_vocab ( model );
return common_tokenize ( vocab , text , add_special , parse_special );
2023-09-28 21:42:38 +02:00
}
2024-10-10 22:57:42 +02:00
std :: vector < llama_token > common_tokenize (
2025-01-12 11:32:42 +02:00
const struct llama_vocab * vocab ,
2023-08-21 23:07:43 +03:00
const std :: string & text ,
2024-04-09 13:44:08 -04:00
bool add_special ,
bool parse_special ) {
2023-08-21 23:07:43 +03:00
// upper limit for the number of tokens
2024-04-09 13:44:08 -04:00
int n_tokens = text . length () + 2 * add_special ;
2023-08-21 23:07:43 +03:00
std :: vector < llama_token > result ( n_tokens );
2025-01-12 11:32:42 +02:00
n_tokens = llama_tokenize ( vocab , text . data (), text . length (), result . data (), result . size (), add_special , parse_special );
2025-06-20 22:13:06 +08:00
if ( n_tokens == std :: numeric_limits < int32_t >:: min ()) {
throw std :: runtime_error ( "Tokenization failed: input text too large, tokenization result exceeds int32_t limit" );
}
2023-08-21 23:07:43 +03:00
if ( n_tokens < 0 ) {
result . resize ( - n_tokens );
2025-01-12 11:32:42 +02:00
int check = llama_tokenize ( vocab , text . data (), text . length (), result . data (), result . size (), add_special , parse_special );
2023-08-21 23:07:43 +03:00
GGML_ASSERT ( check == - n_tokens );
} else {
result . resize ( n_tokens );
}
return result ;
}
2024-10-10 22:57:42 +02:00
std :: string common_token_to_piece ( const struct llama_context * ctx , llama_token token , bool special ) {
2025-01-12 11:32:42 +02:00
const llama_model * model = llama_get_model ( ctx );
const llama_vocab * vocab = llama_model_get_vocab ( model );
return common_token_to_piece ( vocab , token , special );
}
std :: string common_token_to_piece ( const struct llama_vocab * vocab , llama_token token , bool special ) {
2024-07-05 19:01:35 +02:00
std :: string piece ;
piece . resize ( piece . capacity ()); // using string internal cache, 15 bytes + '\n'
2025-01-12 11:32:42 +02:00
const int n_chars = llama_token_to_piece ( vocab , token , & piece [ 0 ], piece . size (), 0 , special );
2024-07-05 19:01:35 +02:00
if ( n_chars < 0 ) {
piece . resize ( - n_chars );
2025-01-12 11:32:42 +02:00
int check = llama_token_to_piece ( vocab , token , & piece [ 0 ], piece . size (), 0 , special );
2024-07-05 19:01:35 +02:00
GGML_ASSERT ( check == - n_chars );
}
else {
piece . resize ( n_chars );
2023-08-21 23:07:43 +03:00
}
2024-07-05 19:01:35 +02:00
return piece ;
2023-08-21 23:07:43 +03:00
}
2023-08-27 14:19:19 +03:00
2025-01-12 11:32:42 +02:00
std :: string common_detokenize ( const struct llama_context * ctx , const std :: vector < llama_token > & tokens , bool special ) {
const llama_model * model = llama_get_model ( ctx );
const llama_vocab * vocab = llama_model_get_vocab ( model );
return common_detokenize ( vocab , tokens , special );
}
std :: string common_detokenize ( const struct llama_vocab * vocab , const std :: vector < llama_token > & tokens , bool special ) {
2024-07-05 19:01:35 +02:00
std :: string text ;
text . resize ( std :: max ( text . capacity (), tokens . size ()));
2025-01-12 11:32:42 +02:00
int32_t n_chars = llama_detokenize ( vocab , tokens . data (), ( int32_t ) tokens . size (), & text [ 0 ], ( int32_t ) text . size (), false , special );
2024-07-05 19:01:35 +02:00
if ( n_chars < 0 ) {
text . resize ( - n_chars );
2025-01-12 11:32:42 +02:00
n_chars = llama_detokenize ( vocab , tokens . data (), ( int32_t ) tokens . size (), & text [ 0 ], ( int32_t ) text . size (), false , special );
2024-07-05 19:01:35 +02:00
GGML_ASSERT ( n_chars <= ( int32_t ) text . size ()); // whitespace trimming is performed after per-token detokenization
2023-08-27 14:19:19 +03:00
}
2024-07-05 19:01:35 +02:00
text . resize ( n_chars );
2023-08-27 14:19:19 +03:00
2023-10-03 09:16:26 +02:00
// NOTE: the original tokenizer decodes bytes after collecting the pieces.
2024-07-05 19:01:35 +02:00
return text ;
2023-08-27 14:19:19 +03:00
}
2023-08-28 17:59:39 +02:00
2024-05-22 20:04:20 +03:00
//
// Embedding utils
//
2024-10-10 22:57:42 +02:00
void common_embd_normalize ( const float * inp , float * out , int n , int embd_norm ) {
2024-03-09 21:27:58 +09:00
double sum = 0.0 ;
2024-06-24 13:30:24 +08:00
switch ( embd_norm ) {
case - 1 : // no normalisation
sum = 1.0 ;
break ;
case 0 : // max absolute
for ( int i = 0 ; i < n ; i ++ ) {
2024-12-18 13:01:41 +02:00
if ( sum < std :: abs ( inp [ i ])) {
sum = std :: abs ( inp [ i ]);
}
2024-06-24 13:30:24 +08:00
}
sum /= 32760.0 ; // make an int16 range
break ;
case 2 : // euclidean
for ( int i = 0 ; i < n ; i ++ ) {
sum += inp [ i ] * inp [ i ];
}
sum = std :: sqrt ( sum );
break ;
default : // p-norm (euclidean is p-norm p=2)
for ( int i = 0 ; i < n ; i ++ ) {
sum += std :: pow ( std :: abs ( inp [ i ]), embd_norm );
}
sum = std :: pow ( sum , 1.0 / embd_norm );
break ;
}
const float norm = sum > 0.0 ? 1.0 / sum : 0.0f ;
2024-03-09 21:27:58 +09:00
for ( int i = 0 ; i < n ; i ++ ) {
out [ i ] = inp [ i ] * norm ;
}
}
2024-10-10 22:57:42 +02:00
float common_embd_similarity_cos ( const float * embd1 , const float * embd2 , int n ){
2024-03-14 10:12:29 +02:00
double sum = 0.0 ;
double sum1 = 0.0 ;
double sum2 = 0.0 ;
for ( int i = 0 ; i < n ; i ++ ) {
sum += embd1 [ i ] * embd2 [ i ];
sum1 += embd1 [ i ] * embd1 [ i ];
sum2 += embd2 [ i ] * embd2 [ i ];
}
2024-06-24 13:30:24 +08:00
// Handle the case where one or both vectors are zero vectors
if ( sum1 == 0.0 || sum2 == 0.0 ) {
if ( sum1 == 0.0 && sum2 == 0.0 ) {
return 1.0f ; // two zero vectors are similar
}
return 0.0f ;
}
2024-03-14 10:12:29 +02:00
return sum / ( sqrt ( sum1 ) * sqrt ( sum2 ));
}
2024-03-15 13:43:02 -07:00
//
// Control vector utils
//
2024-10-10 22:57:42 +02:00
static common_control_vector_data common_control_vector_load_one ( const common_control_vector_load_info & load_info ) {
common_control_vector_data result = { - 1 , {} };
2024-03-15 13:43:02 -07:00
2024-06-27 15:48:07 +01:00
ggml_context * ctx = nullptr ;
struct gguf_init_params meta_gguf_params = {
/* .no_alloc = */ false ,
/* .ctx = */ & ctx ,
};
struct gguf_context * ctx_gguf = gguf_init_from_file ( load_info . fname . c_str (), meta_gguf_params );
if ( ! ctx_gguf ) {
2026-06-28 08:52:15 +03:00
COM_ERR ( "failed to load control vector file from %s \n " , load_info . fname . c_str ());
2024-06-27 15:48:07 +01:00
return result ;
2024-03-15 13:43:02 -07:00
}
2024-06-27 15:48:07 +01:00
int32_t n_tensors = gguf_get_n_tensors ( ctx_gguf );
2024-03-15 13:43:02 -07:00
if ( n_tensors == 0 ) {
2026-06-28 08:52:15 +03:00
COM_WRN ( "no direction tensors found in %s \n " , load_info . fname . c_str ());
2024-03-15 13:43:02 -07:00
}
2024-06-27 15:48:07 +01:00
for ( int i = 0 ; i < n_tensors ; i ++ ) {
std :: string name = gguf_get_tensor_name ( ctx_gguf , i );
2024-03-15 13:43:02 -07:00
2024-06-27 15:48:07 +01:00
int layer_idx = - 1 ;
2024-03-15 13:43:02 -07:00
2024-06-27 15:48:07 +01:00
// split on '.'
size_t dotpos = name . find ( '.' );
if ( dotpos != std :: string :: npos && name . substr ( 0 , dotpos ) == "direction" ) {
try {
layer_idx = std :: stoi ( name . substr ( dotpos + 1 ));
} catch (...) {
layer_idx = - 1 ;
2024-03-15 13:43:02 -07:00
}
}
2024-06-27 15:48:07 +01:00
if ( layer_idx < 0 ) {
2026-06-28 08:52:15 +03:00
COM_ERR ( "invalid/unparsable direction tensor layer index in %s \n " , load_info . fname . c_str ());
2024-06-27 15:48:07 +01:00
result . n_embd = - 1 ;
break ;
} else if ( layer_idx == 0 ) {
2026-06-28 08:52:15 +03:00
COM_ERR ( "invalid (zero) direction tensor layer index in %s \n " , load_info . fname . c_str ());
2024-06-27 15:48:07 +01:00
result . n_embd = - 1 ;
break ;
}
struct ggml_tensor * tensor = ggml_get_tensor ( ctx , name . c_str ());
if ( tensor -> type != GGML_TYPE_F32 ) {
2026-06-28 08:52:15 +03:00
COM_ERR ( "invalid (non-F32) direction tensor type in %s \n " , load_info . fname . c_str ());
2024-06-27 15:48:07 +01:00
result . n_embd = - 1 ;
break ;
}
if ( ggml_n_dims ( tensor ) != 1 ) {
2026-06-28 08:52:15 +03:00
COM_ERR ( "invalid (non-1D) direction tensor shape in %s \n " , load_info . fname . c_str ());
2024-06-27 15:48:07 +01:00
result . n_embd = - 1 ;
break ;
}
if ( result . n_embd == - 1 ) {
result . n_embd = ggml_nelements ( tensor );
} else if ( ggml_nelements ( tensor ) != result . n_embd ) {
2026-06-28 08:52:15 +03:00
COM_ERR ( "direction tensor in %s does not match previous dimensions \n " , load_info . fname . c_str ());
2024-06-27 15:48:07 +01:00
result . n_embd = - 1 ;
break ;
}
// extend if necessary - do not store data for layer 0 (it's not used)
result . data . resize ( std :: max ( result . data . size (), static_cast < size_t > ( result . n_embd * layer_idx )), 0.0f );
const float * src = ( const float * ) tensor -> data ;
float * dst = result . data . data () + result . n_embd * ( layer_idx - 1 ); // layer 1 at [0]
for ( int j = 0 ; j < result . n_embd ; j ++ ) {
dst [ j ] += src [ j ] * load_info . strength ; // allows multiple directions for same layer in same file
}
2024-03-15 13:43:02 -07:00
}
2024-06-27 15:48:07 +01:00
if ( result . n_embd == - 1 ) {
2026-06-28 08:52:15 +03:00
COM_WRN ( "skipping %s due to invalid direction tensors \n " , load_info . fname . c_str ());
2024-06-27 15:48:07 +01:00
result . data . clear ();
}
gguf_free ( ctx_gguf );
ggml_free ( ctx );
2024-03-15 13:43:02 -07:00
return result ;
}
2024-10-10 22:57:42 +02:00
common_control_vector_data common_control_vector_load ( const std :: vector < common_control_vector_load_info > & load_infos ) {
common_control_vector_data result = { - 1 , {} };
2024-03-15 13:43:02 -07:00
for ( const auto & info : load_infos ) {
2024-10-10 22:57:42 +02:00
auto cur = common_control_vector_load_one ( info );
2024-03-15 13:43:02 -07:00
if ( cur . n_embd == - 1 ) {
2024-06-27 15:48:07 +01:00
result . n_embd = - 1 ;
break ;
2024-03-15 13:43:02 -07:00
}
2024-06-27 15:48:07 +01:00
if ( result . n_embd != - 1 && result . n_embd != cur . n_embd ) {
2026-06-28 08:52:15 +03:00
COM_ERR ( "control vectors in %s does not match previous dimensions \n " , info . fname . c_str ());
2024-06-27 15:48:07 +01:00
result . n_embd = - 1 ;
break ;
2024-03-15 13:43:02 -07:00
}
if ( result . n_embd == - 1 ) {
result = std :: move ( cur );
} else {
2024-06-27 15:48:07 +01:00
result . data . resize ( std :: max ( result . data . size (), cur . data . size ()), 0.0f ); // extend if necessary
2024-03-15 13:43:02 -07:00
for ( size_t i = 0 ; i < cur . data . size (); i ++ ) {
result . data [ i ] += cur . data [ i ];
}
}
}
if ( result . n_embd == - 1 ) {
2026-06-28 08:52:15 +03:00
COM_ERR ( "%s" , "no valid control vector files passed \n " );
2024-06-27 15:48:07 +01:00
result . data . clear ();
2024-03-15 13:43:02 -07:00
}
return result ;
}
2025-05-12 14:44:49 +02:00
ggml_opt_dataset_t common_opt_dataset_init ( struct llama_context * ctx , const std :: vector < llama_token > & tokens , int64_t stride ) {
const int64_t ne_datapoint = llama_n_ctx ( ctx );
const int64_t ndata = ( tokens . size () - ne_datapoint - 1 ) / stride ;
ggml_opt_dataset_t result = ggml_opt_dataset_init (
GGML_TYPE_I32 , GGML_TYPE_I32 , ne_datapoint , ne_datapoint , ndata , /*ndata_shard =*/ 1 );
llama_token * data = ( llama_token * ) ggml_opt_dataset_data ( result ) -> data ;
llama_token * labels = ( llama_token * ) ggml_opt_dataset_labels ( result ) -> data ;
for ( int64_t idata = 0 ; idata < ndata ; ++ idata ) {
memcpy ( data + idata * ne_datapoint , tokens . data () + idata * stride + 0 , ne_datapoint * sizeof ( llama_token ));
memcpy ( labels + idata * ne_datapoint , tokens . data () + idata * stride + 1 , ne_datapoint * sizeof ( llama_token ));
}
return result ;
}
2025-08-14 03:03:57 -07:00
ggml_opt_optimizer_params common_opt_lr_pars ( void * userdata ) {
ggml_opt_optimizer_params result = ggml_opt_get_default_optimizer_params ( nullptr );
const lr_opt & d = * ( lr_opt * ) userdata ;
result . adamw . alpha = result . sgd . alpha = d . get_lr ( d . epoch );
result . sgd . wd = result . adamw . wd = d . wd ;
return result ;
}
// TODO make all command line args case-insensitive
static inline bool eq_case_insensitive ( char const * a , char const * b ) {
return !
#if defined(_MSC_VER)
_stricmp
#else
strcasecmp
#endif // defined(_MSC_VER)
( a , b );
}
enum ggml_opt_optimizer_type common_opt_get_optimizer ( const char * n ) {
if ( eq_case_insensitive ( "adamw" , n )) {
return GGML_OPT_OPTIMIZER_TYPE_ADAMW ;
}
if ( eq_case_insensitive ( "sgd" , n )) {
return GGML_OPT_OPTIMIZER_TYPE_SGD ;
}
return GGML_OPT_OPTIMIZER_TYPE_COUNT ;
}
// TODO simplify to use just log and exp
static float const k_log_2 = std :: log ( 2.f );
void lr_opt :: init () {
if ( lr_min > 0 && lr_min < lr0 ) {
float nhalf = std :: log ( lr0 / lr_min ) / k_log_2 ;
float e = epochs ;
if ( decay_epochs > 0 && decay_epochs < e ) {
e = decay_epochs ;
} else {
decay_epochs = e ;
}
scale_epoch = nhalf / e ;
}
}
float lr_opt :: get_lr ( float epoch ) const {
float r = lr_min <= 0 ? lr0 :
epoch >= decay_epochs ? lr_min :
lr0 * std :: pow ( 0.5f , epoch * scale_epoch );
LOG_INF ( "epoch %.2g lr=%.2g \n " , epoch , r );
return r ;
}
2026-02-23 07:04:30 +01:00
bool common_replay_last_token ( struct llama_context * ctx , llama_token last_token , int32_t pos ) {
llama_batch batch = llama_batch_get_one ( & last_token , 1 );
batch . pos = & pos ;
if ( llama_decode ( ctx , batch )) {
LOG_ERR ( "%s: failed to replay last token \n " , __func__ );
return false ;
}
return true ;
}
bool common_prompt_batch_decode (
struct llama_context * ctx ,
2026-06-02 15:44:15 +02:00
const std :: vector < llama_token > & all_tokens ,
int n_new ,
2026-02-23 07:04:30 +01:00
int & n_past ,
int n_batch ,
std :: string_view state_path ,
bool save_state ) {
2026-06-02 15:44:15 +02:00
if ( n_new == 0 ) {
2026-02-23 07:04:30 +01:00
return true ;
}
2026-06-02 15:44:15 +02:00
const int offset = all_tokens . size () - n_new ;
2026-02-23 07:04:30 +01:00
2026-06-02 15:44:15 +02:00
if ( save_state && n_new > 1 ) {
const int n_tokens_before_last = n_new - 1 ;
2026-02-23 07:04:30 +01:00
2026-06-02 15:44:15 +02:00
GGML_ASSERT ( n_new <= n_batch );
2026-02-23 07:04:30 +01:00
// Decode all but the last token so we can save the memory state before decoding the last token.
// This is done so we can restore the session state later and replay the last token.
// Memory implementations in recurrent/hybrid models don't support removing tokens from their
// memory, so we can't just remove the last token from the memory and replay the last token which
// is the reason for this logic.
2026-06-02 15:44:15 +02:00
if ( llama_decode ( ctx , llama_batch_get_one ( const_cast < llama_token *> ( all_tokens . data () + offset ), n_tokens_before_last ))) {
2026-06-28 08:52:15 +03:00
COM_ERR ( "%s" , "failed to eval \n " );
2026-02-23 07:04:30 +01:00
return false ;
}
n_past += n_tokens_before_last ;
2026-06-02 15:44:15 +02:00
llama_state_save_file ( ctx , state_path . data (), all_tokens . data (), all_tokens . size ());
2026-06-28 08:52:15 +03:00
COM_INF ( "saved session before last token to %s, n_new = %zu \n " , state_path . data (), all_tokens . size ());
2026-02-23 07:04:30 +01:00
2026-06-02 15:44:15 +02:00
llama_token last_token = all_tokens . back ();
2026-02-23 07:04:30 +01:00
llama_batch batch = llama_batch_get_one ( & last_token , 1 );
int32_t pos = n_past ;
batch . pos = & pos ;
if ( llama_decode ( ctx , batch )) {
2026-06-28 08:52:15 +03:00
COM_ERR ( "%s" , "failed to eval last token \n " );
2026-02-23 07:04:30 +01:00
return false ;
}
n_past ++ ;
} else {
2026-06-02 15:44:15 +02:00
if ( llama_decode ( ctx , llama_batch_get_one ( const_cast < llama_token *> ( all_tokens . data () + offset ), n_new ))) {
2026-06-28 08:52:15 +03:00
COM_ERR ( "%s" , "failed to eval \n " );
2026-02-23 07:04:30 +01:00
return false ;
}
2026-06-02 15:44:15 +02:00
n_past += n_new ;
2026-02-23 07:04:30 +01:00
}
return true ;
}
2026-05-11 19:09:43 +03:00
size_t common_prompt_checkpoint :: size () const {
2026-06-19 12:08:50 +02:00
return data_tgt . size () + data_dft . size () + data_spec . size ();
2026-05-11 19:09:43 +03:00
}
bool common_prompt_checkpoint :: empty () const {
return data_tgt . empty ();
}
void common_prompt_checkpoint :: clear () {
n_tokens = 0 ;
pos_min = 0 ;
pos_max = 0 ;
data_tgt . clear ();
data_dft . clear ();
2026-06-19 12:08:50 +02:00
data_spec . clear ();
2026-05-11 19:09:43 +03:00
}
void common_prompt_checkpoint :: update_pos (
int64_t n_tokens ,
llama_pos pos_min ,
llama_pos pos_max ) {
this -> n_tokens = n_tokens ;
this -> pos_min = pos_min ;
this -> pos_max = pos_max ;
}
void common_prompt_checkpoint :: update_tgt (
llama_context * ctx ,
llama_seq_id seq_id ,
llama_state_seq_flags flags ) {
if ( ctx == nullptr ) {
return ;
}
const size_t ckpt_size = llama_state_seq_get_size_ext ( ctx , seq_id , flags );
data_tgt . resize ( ckpt_size );
const size_t n = llama_state_seq_get_data_ext ( ctx , data_tgt . data (), ckpt_size , seq_id , flags );
if ( n != ckpt_size ) {
GGML_ABORT ( "checkpoint size mismatch: expected %zu, got %zu \n " , ckpt_size , n );
}
}
void common_prompt_checkpoint :: update_dft (
llama_context * ctx ,
llama_seq_id seq_id ,
llama_state_seq_flags flags ) {
if ( ctx == nullptr ) {
return ;
}
const size_t ckpt_size = llama_state_seq_get_size_ext ( ctx , seq_id , flags );
data_dft . resize ( ckpt_size );
const size_t n = llama_state_seq_get_data_ext ( ctx , data_dft . data (), ckpt_size , seq_id , flags );
if ( n != ckpt_size ) {
GGML_ABORT ( "checkpoint size mismatch: expected %zu, got %zu \n " , ckpt_size , n );
}
}
void common_prompt_checkpoint :: load_tgt (
llama_context * ctx ,
llama_seq_id seq_id ,
llama_state_seq_flags flags ) const {
if ( ctx == nullptr ) {
return ;
}
if ( data_tgt . empty ()) {
return ;
}
const size_t n = llama_state_seq_set_data_ext ( ctx , data_tgt . data (), data_tgt . size (), seq_id , flags );
if ( n != data_tgt . size ()) {
GGML_ABORT ( "checkpoint size mismatch: expected %zu, got %zu \n " , data_tgt . size (), n );
}
}
void common_prompt_checkpoint :: load_dft (
llama_context * ctx ,
llama_seq_id seq_id ,
llama_state_seq_flags flags ) const {
if ( ctx == nullptr ) {
return ;
}
if ( data_dft . empty ()) {
return ;
}
const size_t n = llama_state_seq_set_data_ext ( ctx , data_dft . data (), data_dft . size (), seq_id , flags );
if ( n != data_dft . size ()) {
GGML_ABORT ( "checkpoint size mismatch: expected %zu, got %zu \n " , data_dft . size (), n );
}
}
2026-05-16 20:06:23 +08:00
void common_prompt_checkpoint :: clear_tgt () {
data_tgt . clear ();
}
void common_prompt_checkpoint :: clear_dft () {
data_dft . clear ();
2026-06-19 12:08:50 +02:00
data_spec . clear ();
2026-05-16 20:06:23 +08:00
}