From 2d5c7e2175be033bf3e3c5ec5a60b96f02d10521 Mon Sep 17 00:00:00 2001 From: Ale Date: Wed, 16 Sep 2026 14:48:29 +0200 Subject: [PATCH 1/2] Refactor: Migrate to unified smollm JNI target with KleidiAI support and standard LLMRunner - Consolidate CMake build to a single arm64-v8a target using KleidiAI native dispatch. - Downgrade kotlin serialization to 2.0.0 across modules to match KSP version. - Remove redundant LLM-Runner submodule and implement LLMRunner abstraction locally. - Simplify SmolLM.kt and LLMInference.cpp by delegating core inference loop to LLMRunner. --- app/build.gradle.kts | 2 +- build.gradle.kts | 2 +- hf-model-hub-api/build.gradle.kts | 2 +- smollm/build.gradle.kts | 5 + smollm/src/main/cpp/CMakeLists.txt | 167 +++---- smollm/src/main/cpp/LLMInference.cpp | 359 ++------------- smollm/src/main/cpp/LLMInference.h | 51 +-- smollm/src/main/cpp/LLMRunner.cpp | 415 ++++++++++++++++++ smollm/src/main/cpp/LLMRunner.h | 81 ++++ .../main/java/io/shubham0204/smollm/SmolLM.kt | 92 +--- 10 files changed, 609 insertions(+), 567 deletions(-) create mode 100644 smollm/src/main/cpp/LLMRunner.cpp create mode 100644 smollm/src/main/cpp/LLMRunner.h diff --git a/app/build.gradle.kts b/app/build.gradle.kts index fec4d607..85a39ae5 100644 --- a/app/build.gradle.kts +++ b/app/build.gradle.kts @@ -3,7 +3,7 @@ plugins { alias(libs.plugins.kotlin.android) alias(libs.plugins.kotlin.compose) id("com.google.devtools.ksp") - kotlin("plugin.serialization") version "2.1.0" + kotlin("plugin.serialization") version "2.0.0" } android { diff --git a/build.gradle.kts b/build.gradle.kts index 0da250c3..b804206d 100644 --- a/build.gradle.kts +++ b/build.gradle.kts @@ -6,5 +6,5 @@ plugins { alias(libs.plugins.android.library) apply false id("com.google.devtools.ksp") version "2.0.0-1.0.24" apply false alias(libs.plugins.jetbrains.kotlin.jvm) apply false - kotlin("plugin.serialization") version "2.1.0" apply false + kotlin("plugin.serialization") version "2.0.0" apply false } diff --git a/hf-model-hub-api/build.gradle.kts b/hf-model-hub-api/build.gradle.kts index 6c4a2c05..157ad491 100644 --- a/hf-model-hub-api/build.gradle.kts +++ b/hf-model-hub-api/build.gradle.kts @@ -1,7 +1,7 @@ plugins { id("java-library") alias(libs.plugins.jetbrains.kotlin.jvm) - kotlin("plugin.serialization") version "2.1.0" + kotlin("plugin.serialization") version "2.0.0" } val ktorVersion = "3.0.2" diff --git a/smollm/build.gradle.kts b/smollm/build.gradle.kts index 9c63faa2..91ab7287 100644 --- a/smollm/build.gradle.kts +++ b/smollm/build.gradle.kts @@ -42,12 +42,17 @@ android { arguments += "-DLLAMA_BUILD_COMMON=ON" arguments += "-DLLAMA_CURL=OFF" arguments += "-DGGML_LLAMAFILE=OFF" + arguments += "-DGGML_CPU_KLEIDIAI=ON" + arguments += "-DLLAMA_RUNNER_BACKEND=llama.cpp" // (debugging) uncomment the following line to enable debug builds // and attach hardware-assisted address sanitizer // arguments += "-DCMAKE_BUILD_TYPE=Debug" // arguments += listOf("-DANDROID_SANITIZE=hwaddress") } } + ndk { + abiFilters.add("arm64-v8a") + } } buildTypes { diff --git a/smollm/src/main/cpp/CMakeLists.txt b/smollm/src/main/cpp/CMakeLists.txt index 0531f5c9..6676db3b 100644 --- a/smollm/src/main/cpp/CMakeLists.txt +++ b/smollm/src/main/cpp/CMakeLists.txt @@ -1,26 +1,26 @@ cmake_minimum_required(VERSION 3.22.1) project("smollm") -add_subdirectory(../../../../llama.cpp llama.cpp) +# Configure KleidiAI and Arm architecture for unified single-target build +set(GGML_SYSTEM_ARCH "ARM" CACHE STRING "" FORCE) +set(GGML_CPU_KLEIDIAI ON CACHE BOOL "" FORCE) +set(LLAMA_RUNNER_BACKEND "llama.cpp" CACHE STRING "" FORCE) + set(LLAMA_DIR_RELATIVE "../../../../llama.cpp") get_filename_component(LLAMA_DIR ${LLAMA_DIR_RELATIVE} ABSOLUTE) + +add_subdirectory(../../../../llama.cpp llama.cpp) + set(GGML_DIR ${LLAMA_DIR}/ggml) set(COMMON_DIR ${LLAMA_DIR}/common) set(VENDOR_DIR ${LLAMA_DIR}/vendor) -# -fvisibility=hidden: hide all symbols by default -# -fvisibility-inlines-hidden: hide all inline symbols by default +# Compile options for llama target_compile_options( llama PUBLIC -fvisibility=hidden -fvisibility-inlines-hidden -) -# -ffunction-sections: place each function in its own section -# -fdata-sections: place each data member in its own section -target_compile_options( - llama - PUBLIC -ffunction-sections -fdata-sections ) target_link_options( @@ -30,109 +30,64 @@ target_link_options( -Wl,--exclude-libs,ALL ) -# compiling for different CPU extensions for Arm64 (aarch64) -# See docs/build_arm_flags.md for more details - -function(build_library target_name) - add_library( - ${target_name} - SHARED - LLMInference.cpp - smollm.cpp - ) - target_include_directories( - ${target_name} - PUBLIC - ${COMMON_DIR} - ${GGML_DIR}/include - ${GGML_DIR}/src - ${GGML_DIR}/src/ggml-cpu - ${LLAMA_DIR}/include - ${VENDOR_DIR} - ) - # -fvisibility=hidden: hide all symbols by default - # -fvisibility-inlines-hidden: hide all inline symbols by default - target_compile_options( - ${target_name} - PUBLIC - -fvisibility=hidden -fvisibility-inlines-hidden - ) - # -ffunction-sections: place each function in its own section - # -fdata-sections: place each data member in its own section - target_compile_options( - ${target_name} - PUBLIC - -ffunction-sections -fdata-sections - ) - target_link_libraries( - ${target_name} - android log llama common vulkan - ) - # -Wl,--gc-sections: remove unused sections (garbage collection) - # -flto: link-time optimization - # -Wl,--exclude-libs,ALL: exclude all libraries - target_link_options( - ${target_name} - PRIVATE - -Wl,--gc-sections -flto - -Wl,--exclude-libs,ALL - ) -endfunction() - +# Unified single shared library target using Arm KleidiAI dynamic dispatching +add_library( + smollm + SHARED + LLMRunner.cpp + LLMInference.cpp + smollm.cpp +) -function(build_library_arm64 target_name cpu_flags) - build_library(${target_name}) - set(GGML_SYSTEM_ARCH "ARM") - set(GGML_CPU_KLEIDIAI ON) - set(GGML_OPENMP ON) - target_compile_definitions(${target_name} PRIVATE - GGML_SYSTEM_ARCH=${GGML_SYSTEM_ARCH} - GGML_CPU_KLEIDIAI=$ - GGML_OPENMP=$ - ) - target_compile_options( - ${target_name} - PUBLIC - -DGGML_USE_CPU -DGGML_USE_CPU_AARCH64 ${cpu_flags} -O3 - ) -endfunction() +target_include_directories( + smollm + PUBLIC + ${CMAKE_CURRENT_SOURCE_DIR} + ${COMMON_DIR} + ${GGML_DIR}/include + ${GGML_DIR}/src + ${GGML_DIR}/src/ggml-cpu + ${LLAMA_DIR}/include + ${VENDOR_DIR} +) -function(build_library_armv7a target_name cpu_flags fpu fpu_abi) - build_library(${target_name}) - target_compile_options( - ${target_name} - PUBLIC - -DGGML_USE_CPU ${cpu_flags} ${fpu} ${fpu_abi} -O3 - ) -endfunction() +target_compile_definitions( + smollm + PRIVATE + GGML_SYSTEM_ARCH=${GGML_SYSTEM_ARCH} + GGML_CPU_KLEIDIAI=$ + GGML_USE_CPU + GGML_USE_CPU_AARCH64 +) -function(build_library_universal target_name) - build_library(${target_name}) - target_compile_options( - ${target_name} - PUBLIC - -DGGML_USE_CPU -O3 - ) -endfunction() +target_compile_options( + smollm + PUBLIC + -march=armv8-a + -O3 + -fvisibility=hidden + -fvisibility-inlines-hidden + -ffunction-sections + -fdata-sections +) -build_library_universal("smollm") -if (${ANDROID_ABI} STREQUAL "armeabi-v7a") - build_library_armv7a("smollm_v7a" "-march=armv7-a" "-mfpu=neon-vfpv4" "-mfloat-abi=softfp") -endif() -if (${ANDROID_ABI} STREQUAL "arm64-v8a") - build_library_arm64("smollm_v8" "-march=armv8-a") - # Targets for Arm-v8.2a - build_library_arm64("smollm_v8_2_fp16" "-march=armv8.2-a+fp16") - build_library_arm64("smollm_v8_2_fp16_dotprod" "-march=armv8.2-a+fp16+dotprod") +target_link_libraries( + smollm + android + log + llama + common + vulkan +) - # Targets for Arm-v8.4a - build_library_arm64("smollm_v8_4_fp16_dotprod" "-march=armv8.4-a+fp16+dotprod") - build_library_arm64("smollm_v8_4_fp16_dotprod_sve" "-march=armv8.4-a+fp16+dotprod+sve") - build_library_arm64("smollm_v8_4_fp16_dotprod_i8mm" "-march=armv8.4-a+fp16+dotprod+i8mm") - build_library_arm64("smollm_v8_4_fp16_dotprod_i8mm_sve" "-march=armv8.4-a+fp16+dotprod+i8mm+sve") -endif() +target_link_options( + smollm + PRIVATE + -Wl,--gc-sections -flto + -Wl,--exclude-libs,ALL +) -# library target for GGUFReader +# Library target for GGUFReader set(TARGET_NAME_GGUF_READER ggufreader) add_library(${TARGET_NAME_GGUF_READER} SHARED GGUFReader.cpp) target_include_directories( diff --git a/smollm/src/main/cpp/LLMInference.cpp b/smollm/src/main/cpp/LLMInference.cpp index bac714cb..15af1535 100644 --- a/smollm/src/main/cpp/LLMInference.cpp +++ b/smollm/src/main/cpp/LLMInference.cpp @@ -1,349 +1,52 @@ #include "LLMInference.h" -#include -#include -#include -#include +#include -#define TAG "[SmolLMAndroid-Cpp]" -#define LOGi(...) __android_log_print(ANDROID_LOG_INFO, TAG, __VA_ARGS__) -#define LOGe(...) __android_log_print(ANDROID_LOG_ERROR, TAG, __VA_ARGS__) +LLMInference::LLMInference() : m_runner(std::make_unique()) {} -void -LLMInference::loadModel(const char *model_path, float minP, float temperature, bool storeChats, long contextSize, - const char *chatTemplate, int nThreads, bool useMmap, bool useMlock) { - LOGi("loading model with" - "\n\tmodel_path = %s" - "\n\tminP = %f" - "\n\ttemperature = %f" - "\n\tstoreChats = %d" - "\n\tcontextSize = %li" - "\n\tchatTemplate = %s" - "\n\tnThreads = %d" - "\n\tuseMmap = %d" - "\n\tuseMlock = %d", - model_path, minP, temperature, storeChats, contextSize, chatTemplate, nThreads, useMmap, useMlock); +LLMInference::~LLMInference() = default; - // load dynamic backends - ggml_backend_load_all(); +void LLMInference::loadModel(const char* modelPath, float minP, float temperature, bool storeChats, + long contextSize, const char* chatTemplate, int nThreads, + bool useMmap, bool useMlock) { + smollm::RunnerParams params; + params.minP = minP; + params.temperature = temperature; + params.storeChats = storeChats; + params.contextSize = contextSize; + params.chatTemplate = chatTemplate ? chatTemplate : ""; + params.nThreads = nThreads; + params.useMmap = useMmap; + params.useMlock = useMlock; - // create an instance of llama_model - llama_model_params model_params = llama_model_default_params(); - model_params.use_mmap = useMmap; - model_params.use_mlock = useMlock; - _model = llama_model_load_from_file(model_path, model_params); - if (!_model) { - LOGe("failed to load model from %s", model_path); - throw std::runtime_error("loadModel() failed"); + if (!m_runner->load_model(modelPath, params)) { + throw std::runtime_error("Runner::load_model() failed to load model from " + std::string(modelPath)); } - - // create an instance of llama_context - llama_context_params ctx_params = llama_context_default_params(); - ctx_params.n_ctx = contextSize; - ctx_params.n_batch = contextSize; - ctx_params.n_threads = nThreads; - ctx_params.no_perf = true; // disable performance metrics - _ctx = llama_init_from_model(_model, ctx_params); - if (!_ctx) { - LOGe("llama_new_context_with_model() returned null)"); - throw std::runtime_error("llama_new_context_with_model() returned null"); - } - - // create an instance of llama_sampler - llama_sampler_chain_params sampler_params = llama_sampler_chain_default_params(); - sampler_params.no_perf = true; // disable performance metrics - _sampler = llama_sampler_chain_init(sampler_params); - llama_sampler_chain_add(_sampler, llama_sampler_init_temp(temperature)); - llama_sampler_chain_add(_sampler, llama_sampler_init_dist(LLAMA_DEFAULT_SEED)); - - _formattedMessages = std::vector(llama_n_ctx(_ctx)); - _messages.clear(); - - if (chatTemplate == nullptr) { - _chatTemplate = llama_model_chat_template(_model, nullptr); - } else { - _chatTemplate = strdup(chatTemplate); - } - this->_storeChats = storeChats; -} - -void -LLMInference::addChatMessage(const char *message, const char *role) { - _messages.push_back({strdup(role), strdup(message)}); } -float -LLMInference::getResponseGenerationTime() const { - return (float) _responseNumTokens / (_responseGenerationTime / 1e6); +void LLMInference::addChatMessage(const char* message, const char* role) { + m_runner->add_chat_message(role, message); } -int -LLMInference::getContextSizeUsed() const { - return _nCtxUsed; +float LLMInference::getResponseGenerationTime() const { + return m_runner->get_tokens_per_second(); } -bool -LLMInference::startCompletion(const char *query) { - if (!_storeChats) { - _formattedMessages.clear(); - _formattedMessages = std::vector(llama_n_ctx(_ctx)); - } - _responseGenerationTime = 0; - _responseNumTokens = 0; - addChatMessage(query, "user"); - // apply the chat-template - std::vector messages; - for (const llama_chat_message& message : _messages) { - common_chat_msg msg; - msg.role = message.role; - msg.content = message.content; - messages.push_back(msg); - } - auto templates = common_chat_templates_init(_model, _chatTemplate ? _chatTemplate : ""); - - common_chat_templates_inputs inputs; - inputs.messages = messages; - - // Try Jinja rendering first with tools defined to prevent "tojson on Undefined" errors. - // If Jinja fails (e.g. unsupported filters like lstrip), fall back to legacy rendering. - inputs.use_jinja = true; - inputs.chat_template_kwargs["tools"] = "[]"; - - std::string prompt; - bool usedJinja = true; - try { - prompt = common_chat_templates_apply(templates.get(), inputs).prompt; - } catch (const std::exception &e) { - LOGe("Jinja template failed: %s — retrying with legacy renderer", e.what()); - inputs.use_jinja = false; - inputs.chat_template_kwargs.clear(); - prompt = common_chat_templates_apply(templates.get(), inputs).prompt; - usedJinja = false; - } - _promptTokens = common_tokenize(llama_model_get_vocab(_model), prompt, true, true); - - // create a llama_batch containing a single sequence - // see llama_batch_init for more details - _batch = new llama_batch(); - _batch->token = _promptTokens.data(); - _batch->n_tokens = _promptTokens.size(); - - return usedJinja; +int LLMInference::getContextSizeUsed() const { + return m_runner->get_context_size_used(); } -// taken from: -// https://github.com/ggerganov/llama.cpp/blob/master/examples/llama.android/llama/src/main/cpp/llama-android.cpp#L38 -bool -LLMInference::_isValidUtf8(const char *response) { - if (!response) { - return true; - } - const unsigned char *bytes = (const unsigned char *) response; - int num; - while (*bytes != 0x00) { - if ((*bytes & 0x80) == 0x00) { - // U+0000 to U+007F - num = 1; - } else if ((*bytes & 0xE0) == 0xC0) { - // U+0080 to U+07FF - num = 2; - } else if ((*bytes & 0xF0) == 0xE0) { - // U+0800 to U+FFFF - num = 3; - } else if ((*bytes & 0xF8) == 0xF0) { - // U+10000 to U+10FFFF - num = 4; - } else { - return false; - } - - bytes += 1; - for (int i = 1; i < num; ++i) { - if ((*bytes & 0xC0) != 0x80) { - return false; - } - bytes += 1; - } - } - return true; +bool LLMInference::startCompletion(const char* query) { + return m_runner->start_completion(query); } -std::string -LLMInference::completionLoop() { - // check if the length of the inputs to the model - // have exceeded the context size of the model - uint32_t contextSize = llama_n_ctx(_ctx); - _nCtxUsed = llama_memory_seq_pos_max(llama_get_memory(_ctx), 0) + 1; - if (_nCtxUsed + _batch->n_tokens > contextSize) { - throw std::runtime_error("context size reached"); - } - - auto start = ggml_time_us(); - // run the model - if (llama_decode(_ctx, *_batch) < 0) { - throw std::runtime_error("llama_decode() failed"); - } - - // sample a token and check if it is an EOG (end of generation token) - // convert the integer token to its corresponding word-piece - _currToken = llama_sampler_sample(_sampler, _ctx, -1); - if (llama_vocab_is_eog(llama_model_get_vocab(_model), _currToken)) { - addChatMessage(strdup(_response.data()), "assistant"); - _response.clear(); - return "[EOG]"; - } - std::string piece = common_token_to_piece(_ctx, _currToken, true); - auto end = ggml_time_us(); - _responseGenerationTime += (end - start); - _responseNumTokens += 1; - _cacheResponseTokens += piece; - - // re-init the batch with the newly predicted token - // key, value pairs of all previous tokens have been cached - // in the KV cache - _batch->token = &_currToken; - _batch->n_tokens = 1; - - if (_isValidUtf8(_cacheResponseTokens.c_str())) { - _response += _cacheResponseTokens; - std::string valid_utf8_piece = _cacheResponseTokens; - _cacheResponseTokens.clear(); - return valid_utf8_piece; - } - - return ""; +std::string LLMInference::completionLoop() { + return m_runner->completion_loop(); } -void -LLMInference::stopCompletion() { - if (_storeChats) { - addChatMessage(_response.c_str(), "assistant"); - } - _response.clear(); +void LLMInference::stopCompletion() { + m_runner->stop_completion(); } -LLMInference::~LLMInference() { - // free memory held by the message text in messages - // (as we had used strdup() to create a malloc'ed copy) - for (llama_chat_message &message: _messages) { - free(const_cast(message.role)); - free(const_cast(message.content)); - } - llama_free(_ctx); - llama_model_free(_model); - delete _batch; - llama_sampler_free(_sampler); -} - -std::string -LLMInference::benchModel(int pp, int tg, int pl, int nr) { - g_batch = llama_batch_init(pp, 0, pl); - auto pp_avg = 0.0; - auto tg_avg = 0.0; - auto pp_std = 0.0; - auto tg_std = 0.0; - - const uint32_t n_ctx = llama_n_ctx(this->_ctx); - LOGi("n_ctx = %d", n_ctx); - - int i, j; - int nri; - for (nri = 0; nri < nr; nri++) { - LOGi("Benchmark prompt processing (pp = %d)", pp); - - common_batch_clear(g_batch); - - const int n_tokens = pp; - for (i = 0; i < n_tokens; i++) { - common_batch_add(g_batch, 1, i, { 0 }, false); - } - - g_batch.logits[g_batch.n_tokens - 1] = true; - llama_memory_clear(llama_get_memory(this->_ctx), false); - - const auto t_pp_start = ggml_time_us(); - if (llama_decode(this->_ctx, g_batch) != 0) { - LOGe("llama_decode() failed during prompt processing"); - } - const auto t_pp_end = ggml_time_us(); - - // bench text generation - - LOGi("Benchmark text generation (tg = %d)", tg); - - llama_memory_clear(llama_get_memory(this->_ctx), false); - const auto t_tg_start = ggml_time_us(); - for (i = 0; i < tg; i++) { - common_batch_clear(g_batch); - for (j = 0; j < pl; j++) { - common_batch_add(g_batch, 0, i, { j }, true); - } - - if (llama_decode(this->_ctx, g_batch) != 0) { - LOGe("llama_decode() failed during text generation"); - } - } - const auto t_tg_end = ggml_time_us(); - - llama_memory_clear(llama_get_memory(this->_ctx), false); - - const auto t_pp = double(t_pp_end - t_pp_start) / 1000000.0; - const auto t_tg = double(t_tg_end - t_tg_start) / 1000000.0; - - const auto speed_pp = double(pp) / t_pp; - const auto speed_tg = double(pl * tg) / t_tg; - - pp_avg += speed_pp; - tg_avg += speed_tg; - - pp_std += speed_pp * speed_pp; - tg_std += speed_tg * speed_tg; - - LOGi("pp %f t/s, tg %f t/s", speed_pp, speed_tg); - } - - llama_batch_free(g_batch); - - pp_avg /= double(nr); - tg_avg /= double(nr); - - if (nr > 1) { - pp_std = sqrt(pp_std / double(nr - 1) - pp_avg * pp_avg * double(nr) / double(nr - 1)); - tg_std = sqrt(tg_std / double(nr - 1) - tg_avg * tg_avg * double(nr) / double(nr - 1)); - } else { - pp_std = 0; - tg_std = 0; - } - - char model_desc[128]; - llama_model_desc(this->_model, model_desc, sizeof(model_desc)); - - const auto model_size = double(llama_model_size(this->_model)) / 1024.0 / 1024.0 / 1024.0; - const auto model_n_params = double(llama_model_n_params(this->_model)) / 1e9; - - std::vector backends; - for (size_t i = 0; i < ggml_backend_reg_count(); i++) { - auto* reg = ggml_backend_reg_get(i); - std::string name = ggml_backend_reg_name(reg); - if (name != "CPU") { - backends.push_back(ggml_backend_reg_name(reg)); - } - } - std::ostringstream str; - for (size_t i = 0; i < backends.size(); i++) { - str << backends[i]; - if (i < backends.size() - 1) { - str << ","; - } - } - const auto backend = str.str(); - - std::stringstream result; - result << std::setprecision(3); - result << "| model | size | params | backend | test | t/s |\n"; - result << "| --- | --- | --- | --- | --- | --- |\n"; - result << "| " << model_desc << " | " << model_size << "GiB | " << model_n_params << "B | " << backend << " | pp " - << pp << " | " << pp_avg << " ± " << pp_std << " |\n"; - result << "| " << model_desc << " | " << model_size << "GiB | " << model_n_params << "B | " << backend << " | tg " - << tg << " | " << tg_avg << " ± " << tg_std << " |\n"; - return result.str(); +std::string LLMInference::benchModel(int pp, int tg, int pl, int nr) { + return m_runner->bench_model(pp, tg, pl, nr); } diff --git a/smollm/src/main/cpp/LLMInference.h b/smollm/src/main/cpp/LLMInference.h index dfa37bd7..513bf5d2 100644 --- a/smollm/src/main/cpp/LLMInference.h +++ b/smollm/src/main/cpp/LLMInference.h @@ -1,46 +1,17 @@ #pragma once -#include "chat.h" -#include "common.h" -#include "llama.h" + +#include "LLMRunner.h" +#include #include -#include class LLMInference { - // llama.cpp-specific types - llama_context* _ctx; - llama_model* _model; - llama_sampler* _sampler; - llama_token _currToken; - llama_batch* _batch; - - llama_batch g_batch; - - // container to store user/assistant messages in the chat - std::vector _messages; - // stores the string generated after applying - // the chat-template to all messages in `_messages` - std::vector _formattedMessages; - // stores the tokens for the last query - // appended to `_messages` - std::vector _promptTokens; - const char* _chatTemplate; - - // stores the complete response for the given query - std::string _response; - std::string _cacheResponseTokens; - // whether to cache previous messages in `_messages` - bool _storeChats; - - // response generation metrics - int64_t _responseGenerationTime = 0; - long _responseNumTokens = 0; - - // length of context window consumed during the conversation - int _nCtxUsed = 0; - - bool _isValidUtf8(const char* response); - - public: +private: + std::unique_ptr m_runner; + +public: + LLMInference(); + ~LLMInference(); + void loadModel(const char* modelPath, float minP, float temperature, bool storeChats, long contextSize, const char* chatTemplate, int nThreads, bool useMmap, bool useMlock); @@ -59,5 +30,5 @@ class LLMInference { void stopCompletion(); - ~LLMInference(); + smollm::LLMRunner* getRunner() { return m_runner.get(); } }; \ No newline at end of file diff --git a/smollm/src/main/cpp/LLMRunner.cpp b/smollm/src/main/cpp/LLMRunner.cpp new file mode 100644 index 00000000..26622d7b --- /dev/null +++ b/smollm/src/main/cpp/LLMRunner.cpp @@ -0,0 +1,415 @@ +#include "LLMRunner.h" + +#include +#include +#include +#include +#include +#include + +#define TAG "[SmolLM-LLMRunner]" +#define LOGi(...) __android_log_print(ANDROID_LOG_INFO, TAG, __VA_ARGS__) +#define LOGe(...) __android_log_print(ANDROID_LOG_ERROR, TAG, __VA_ARGS__) + +namespace smollm { + +LLMRunner::LLMRunner() = default; + +LLMRunner::~LLMRunner() { + free_resources(); +} + +void LLMRunner::free_resources() { + for (auto& msg : m_messages) { + free(const_cast(msg.role)); + free(const_cast(msg.content)); + } + m_messages.clear(); + + if (m_step_batch) { + delete m_step_batch; + m_step_batch = nullptr; + } + if (m_sampler) { + llama_sampler_free(m_sampler); + m_sampler = nullptr; + } + if (m_ctx) { + llama_free(m_ctx); + m_ctx = nullptr; + } + if (m_model) { + llama_model_free(m_model); + m_model = nullptr; + } +} + +bool LLMRunner::load_model(const std::string& model_path, const RunnerParams& params) { + LOGi("Runner::load_model loading: %s (threads=%d, ctx=%ld, mmap=%d)", + model_path.c_str(), params.nThreads, params.contextSize, params.useMmap); + + free_resources(); + m_params = params; + + // Initialize ggml dynamic backends (CPU/KleidiAI, Vulkan if available) + ggml_backend_load_all(); + + llama_model_params model_params = llama_model_default_params(); + model_params.use_mmap = params.useMmap; + model_params.use_mlock = params.useMlock; + + m_model = llama_model_load_from_file(model_path.c_str(), model_params); + if (!m_model) { + LOGe("Runner::load_model failed to load model from %s", model_path.c_str()); + return false; + } + + llama_context_params ctx_params = llama_context_default_params(); + ctx_params.n_ctx = params.contextSize; + ctx_params.n_batch = params.contextSize; + ctx_params.n_threads = params.nThreads; + ctx_params.no_perf = true; + + m_ctx = llama_init_from_model(m_model, ctx_params); + if (!m_ctx) { + LOGe("Runner::load_model llama_init_from_model returned null"); + llama_model_free(m_model); + m_model = nullptr; + return false; + } + + llama_sampler_chain_params sampler_params = llama_sampler_chain_default_params(); + sampler_params.no_perf = true; + m_sampler = llama_sampler_chain_init(sampler_params); + llama_sampler_chain_add(m_sampler, llama_sampler_init_temp(params.temperature)); + llama_sampler_chain_add(m_sampler, llama_sampler_init_dist(LLAMA_DEFAULT_SEED)); + + if (params.chatTemplate.empty()) { + const char* tmpl = llama_model_chat_template(m_model, nullptr); + m_chat_template = tmpl ? tmpl : ""; + } else { + m_chat_template = params.chatTemplate; + } + + LOGi("Runner::load_model success"); + return true; +} + +std::vector LLMRunner::tokenize(const std::string& prompt, bool add_special, bool parse_special) { + if (!m_model) { + LOGe("Runner::tokenize called with null model"); + return {}; + } + return common_tokenize(llama_model_get_vocab(m_model), prompt, add_special, parse_special); +} + +void LLMRunner::add_chat_message(const std::string& role, const std::string& message) { + m_messages.push_back({strdup(role.c_str()), strdup(message.c_str())}); +} + +std::pair LLMRunner::format_chat_prompt(const std::string& user_query) { + add_chat_message("user", user_query); + + std::vector messages; + for (const auto& msg : m_messages) { + common_chat_msg cmsg; + cmsg.role = msg.role; + cmsg.content = msg.content; + messages.push_back(cmsg); + } + + auto templates = common_chat_templates_init(m_model, m_chat_template.c_str()); + common_chat_templates_inputs inputs; + inputs.messages = messages; + inputs.use_jinja = true; + inputs.chat_template_kwargs["tools"] = "[]"; + + std::string prompt; + bool used_jinja = true; + try { + prompt = common_chat_templates_apply(templates.get(), inputs).prompt; + } catch (const std::exception& e) { + LOGi("Jinja template formatting failed: %s, falling back to legacy", e.what()); + inputs.use_jinja = false; + inputs.chat_template_kwargs.clear(); + prompt = common_chat_templates_apply(templates.get(), inputs).prompt; + used_jinja = false; + } + + return {prompt, used_jinja}; +} + +bool LLMRunner::generate(const std::vector& tokens, TokenCallback callback_stream) { + if (!m_ctx || !m_model || !m_sampler || tokens.empty()) { + return false; + } + + m_generation_time_us = 0; + m_generated_tokens_count = 0; + m_accumulated_response.clear(); + m_utf8_token_cache.clear(); + + llama_batch batch = llama_batch_init(tokens.size(), 0, 1); + for (size_t i = 0; i < tokens.size(); ++i) { + common_batch_add(batch, tokens[i], i, {0}, false); + } + batch.logits[batch.n_tokens - 1] = true; + + if (llama_decode(m_ctx, batch) != 0) { + LOGe("Runner::generate prompt evaluation failed"); + llama_batch_free(batch); + return false; + } + llama_batch_free(batch); + + llama_batch step_batch = llama_batch_init(1, 0, 1); + const uint32_t context_size = llama_n_ctx(m_ctx); + bool should_continue = true; + + while (should_continue) { + m_n_ctx_used = llama_memory_seq_pos_max(llama_get_memory(m_ctx), 0) + 1; + if (m_n_ctx_used >= context_size) { + LOGi("Context size limit reached: %d >= %u", m_n_ctx_used, context_size); + break; + } + + auto t0 = ggml_time_us(); + llama_token token = llama_sampler_sample(m_sampler, m_ctx, -1); + + if (llama_vocab_is_eog(llama_model_get_vocab(m_model), token)) { + break; + } + + std::string piece = common_token_to_piece(m_ctx, token, true); + auto t1 = ggml_time_us(); + m_generation_time_us += (t1 - t0); + m_generated_tokens_count++; + + m_utf8_token_cache += piece; + if (is_valid_utf8(m_utf8_token_cache.c_str())) { + m_accumulated_response += m_utf8_token_cache; + if (callback_stream) { + should_continue = callback_stream(m_utf8_token_cache); + } + m_utf8_token_cache.clear(); + } + + if (!should_continue) { + break; + } + + common_batch_clear(step_batch); + common_batch_add(step_batch, token, m_n_ctx_used, {0}, true); + if (llama_decode(m_ctx, step_batch) != 0) { + LOGe("llama_decode step failed"); + break; + } + } + + llama_batch_free(step_batch); + + if (m_params.storeChats && !m_accumulated_response.empty()) { + add_chat_message("assistant", m_accumulated_response); + } + + return true; +} + +bool LLMRunner::start_completion(const std::string& query) { + if (!m_params.storeChats) { + for (auto& msg : m_messages) { + free(const_cast(msg.role)); + free(const_cast(msg.content)); + } + m_messages.clear(); + } + + m_generation_time_us = 0; + m_generated_tokens_count = 0; + m_accumulated_response.clear(); + m_utf8_token_cache.clear(); + + auto [prompt, used_jinja] = format_chat_prompt(query); + m_prompt_tokens = tokenize(prompt, true, true); + + if (m_step_batch) { + delete m_step_batch; + } + m_step_batch = new llama_batch(); + m_step_batch->token = m_prompt_tokens.data(); + m_step_batch->n_tokens = m_prompt_tokens.size(); + + return used_jinja; +} + +std::string LLMRunner::completion_loop() { + if (!m_ctx || !m_model || !m_step_batch) { + throw std::runtime_error("Runner not initialized for completion"); + } + + uint32_t context_size = llama_n_ctx(m_ctx); + m_n_ctx_used = llama_memory_seq_pos_max(llama_get_memory(m_ctx), 0) + 1; + if (m_n_ctx_used + m_step_batch->n_tokens > context_size) { + throw std::runtime_error("Context size limit reached"); + } + + auto start = ggml_time_us(); + if (llama_decode(m_ctx, *m_step_batch) < 0) { + throw std::runtime_error("llama_decode() failed"); + } + + m_curr_token = llama_sampler_sample(m_sampler, m_ctx, -1); + if (llama_vocab_is_eog(llama_model_get_vocab(m_model), m_curr_token)) { + if (m_params.storeChats && !m_accumulated_response.empty()) { + add_chat_message("assistant", m_accumulated_response); + } + m_accumulated_response.clear(); + return "[EOG]"; + } + + std::string piece = common_token_to_piece(m_ctx, m_curr_token, true); + auto end = ggml_time_us(); + m_generation_time_us += (end - start); + m_generated_tokens_count += 1; + m_utf8_token_cache += piece; + + m_step_batch->token = &m_curr_token; + m_step_batch->n_tokens = 1; + + if (is_valid_utf8(m_utf8_token_cache.c_str())) { + m_accumulated_response += m_utf8_token_cache; + std::string valid_piece = m_utf8_token_cache; + m_utf8_token_cache.clear(); + return valid_piece; + } + + return ""; +} + +void LLMRunner::stop_completion() { + if (m_params.storeChats && !m_accumulated_response.empty()) { + add_chat_message("assistant", m_accumulated_response); + } + m_accumulated_response.clear(); +} + +float LLMRunner::get_tokens_per_second() const { + if (m_generation_time_us <= 0) return 0.0f; + return static_cast(m_generated_tokens_count) / (static_cast(m_generation_time_us) / 1e6f); +} + +int LLMRunner::get_context_size_used() const { + return m_n_ctx_used; +} + +bool LLMRunner::is_valid_utf8(const char* str) const { + if (!str) return true; + const auto* bytes = reinterpret_cast(str); + while (*bytes != 0x00) { + int num = 0; + if ((*bytes & 0x80) == 0x00) num = 1; + else if ((*bytes & 0xE0) == 0xC0) num = 2; + else if ((*bytes & 0xF0) == 0xE0) num = 3; + else if ((*bytes & 0xF8) == 0xF0) num = 4; + else return false; + + bytes += 1; + for (int i = 1; i < num; ++i) { + if ((*bytes & 0xC0) != 0x80) return false; + bytes += 1; + } + } + return true; +} + +std::string LLMRunner::bench_model(int pp, int tg, int pl, int nr) { + llama_batch g_batch = llama_batch_init(pp, 0, pl); + auto pp_avg = 0.0; + auto tg_avg = 0.0; + auto pp_std = 0.0; + auto tg_std = 0.0; + + const uint32_t n_ctx = llama_n_ctx(m_ctx); + LOGi("bench_model: n_ctx = %u", n_ctx); + + for (int nri = 0; nri < nr; nri++) { + common_batch_clear(g_batch); + for (int i = 0; i < pp; i++) { + common_batch_add(g_batch, 1, i, {0}, false); + } + g_batch.logits[g_batch.n_tokens - 1] = true; + llama_memory_clear(llama_get_memory(m_ctx), false); + + const auto t_pp_start = ggml_time_us(); + llama_decode(m_ctx, g_batch); + const auto t_pp_end = ggml_time_us(); + + llama_memory_clear(llama_get_memory(m_ctx), false); + const auto t_tg_start = ggml_time_us(); + for (int i = 0; i < tg; i++) { + common_batch_clear(g_batch); + for (int j = 0; j < pl; j++) { + common_batch_add(g_batch, 0, i, {j}, true); + } + llama_decode(m_ctx, g_batch); + } + const auto t_tg_end = ggml_time_us(); + + llama_memory_clear(llama_get_memory(m_ctx), false); + + const auto t_pp = double(t_pp_end - t_pp_start) / 1000000.0; + const auto t_tg = double(t_tg_end - t_tg_start) / 1000000.0; + const auto speed_pp = double(pp) / t_pp; + const auto speed_tg = double(pl * tg) / t_tg; + + pp_avg += speed_pp; + tg_avg += speed_tg; + pp_std += speed_pp * speed_pp; + tg_std += speed_tg * speed_tg; + } + + llama_batch_free(g_batch); + + pp_avg /= double(nr); + tg_avg /= double(nr); + if (nr > 1) { + pp_std = sqrt(pp_std / double(nr - 1) - pp_avg * pp_avg * double(nr) / double(nr - 1)); + tg_std = sqrt(tg_std / double(nr - 1) - tg_avg * tg_avg * double(nr) / double(nr - 1)); + } else { + pp_std = 0; + tg_std = 0; + } + + char model_desc[128]; + llama_model_desc(m_model, model_desc, sizeof(model_desc)); + + const auto model_size = double(llama_model_size(m_model)) / 1024.0 / 1024.0 / 1024.0; + const auto model_n_params = double(llama_model_n_params(m_model)) / 1e9; + + std::vector backends; + for (size_t i = 0; i < ggml_backend_reg_count(); i++) { + auto* reg = ggml_backend_reg_get(i); + std::string name = ggml_backend_reg_name(reg); + if (name != "CPU") { + backends.push_back(name); + } + } + std::ostringstream str; + for (size_t i = 0; i < backends.size(); i++) { + str << backends[i]; + if (i < backends.size() - 1) str << ","; + } + + std::stringstream result; + result << std::setprecision(3); + result << "| model | size | params | backend | test | t/s |\n"; + result << "| --- | --- | --- | --- | --- | --- |\n"; + result << "| " << model_desc << " | " << model_size << "GiB | " << model_n_params << "B | " << str.str() << " | pp " + << pp << " | " << pp_avg << " ± " << pp_std << " |\n"; + result << "| " << model_desc << " | " << model_size << "GiB | " << model_n_params << "B | " << str.str() << " | tg " + << tg << " | " << tg_avg << " ± " << tg_std << " |\n"; + return result.str(); +} + +} // namespace smollm + diff --git a/smollm/src/main/cpp/LLMRunner.h b/smollm/src/main/cpp/LLMRunner.h new file mode 100644 index 00000000..97bd06c9 --- /dev/null +++ b/smollm/src/main/cpp/LLMRunner.h @@ -0,0 +1,81 @@ +#pragma once + +#include "chat.h" +#include "common.h" +#include "llama.h" + +#include +#include +#include +#include + +namespace smollm { + +struct RunnerParams { + float minP = 0.1f; + float temperature = 0.8f; + bool storeChats = true; + long contextSize = 1024; + std::string chatTemplate; + int nThreads = 4; + bool useMmap = true; + bool useMlock = false; +}; + +class LLMRunner { +public: + using TokenCallback = std::function; + + LLMRunner(); + ~LLMRunner(); + + // Standardized Runner lifecycle: + // 1. Runner::load_model(model_path, params) + bool load_model(const std::string& model_path, const RunnerParams& params); + + // 2. Runner::tokenize(prompt) + std::vector tokenize(const std::string& prompt, bool add_special = true, bool parse_special = true); + + // 3. Runner::generate(tokens, callback_stream) + bool generate(const std::vector& tokens, TokenCallback callback_stream); + + // Conversation management + void add_chat_message(const std::string& role, const std::string& message); + std::pair format_chat_prompt(const std::string& user_query); + + // Step-by-step completion methods for existing JNI interface + bool start_completion(const std::string& query); + std::string completion_loop(); + void stop_completion(); + + // Metrics and benchmarking + float get_tokens_per_second() const; + int get_context_size_used() const; + std::string bench_model(int pp, int tg, int pl, int nr); + +private: + bool is_valid_utf8(const char* str) const; + void free_resources(); + + llama_model* m_model = nullptr; + llama_context* m_ctx = nullptr; + llama_sampler* m_sampler = nullptr; + + RunnerParams m_params; + std::string m_chat_template; + + std::vector m_messages; + std::vector m_prompt_tokens; + llama_batch* m_step_batch = nullptr; + llama_token m_curr_token = 0; + + std::string m_accumulated_response; + std::string m_utf8_token_cache; + + int64_t m_generation_time_us = 0; + long m_generated_tokens_count = 0; + int m_n_ctx_used = 0; +}; + +} // namespace smollm + diff --git a/smollm/src/main/java/io/shubham0204/smollm/SmolLM.kt b/smollm/src/main/java/io/shubham0204/smollm/SmolLM.kt index d2c7c25f..e87e88cb 100644 --- a/smollm/src/main/java/io/shubham0204/smollm/SmolLM.kt +++ b/smollm/src/main/java/io/shubham0204/smollm/SmolLM.kt @@ -16,106 +16,18 @@ package io.shubham0204.smollm -import android.os.Build -import android.util.Log import kotlinx.coroutines.Dispatchers import kotlinx.coroutines.flow.Flow import kotlinx.coroutines.flow.flow import kotlinx.coroutines.withContext -import java.io.File -import java.io.FileNotFoundException /** This class interacts with the JNI binding and provides a Kotlin API to infer a GGUF LLM model */ class SmolLM { companion object { init { - val logTag = SmolLM::class.java.simpleName - - // check if the following CPU features are available, - // and load the native library accordingly - val cpuFeatures = getCPUFeatures() - val hasFp16 = cpuFeatures.contains("fp16") || cpuFeatures.contains("fphp") - val hasDotProd = cpuFeatures.contains("dotprod") || cpuFeatures.contains("asimddp") - val hasSve = cpuFeatures.contains("sve") - val hasI8mm = cpuFeatures.contains("i8mm") - val isAtLeastArmV82 = - cpuFeatures.contains("asimd") && - cpuFeatures.contains("crc32") && - cpuFeatures.contains("aes") - val isAtLeastArmV84 = cpuFeatures.contains("dcpop") && cpuFeatures.contains("uscat") - - Log.d(logTag, "CPU features: $cpuFeatures") - Log.d(logTag, "- hasFp16: $hasFp16") - Log.d(logTag, "- hasDotProd: $hasDotProd") - Log.d(logTag, "- hasSve: $hasSve") - Log.d(logTag, "- hasI8mm: $hasI8mm") - Log.d(logTag, "- isAtLeastArmV82: $isAtLeastArmV82") - Log.d(logTag, "- isAtLeastArmV84: $isAtLeastArmV84") - - // Check if the app is running in an emulated device - // Note, this is not the OFFICIAL way to check if the app is running - // on an emulator - val isEmulated = - (Build.HARDWARE.contains("goldfish") || Build.HARDWARE.contains("ranchu")) - Log.d(logTag, "isEmulated: $isEmulated") - - if (!isEmulated) { - if (supportsArm64V8a()) { - if (isAtLeastArmV84 && hasSve && hasI8mm && hasFp16 && hasDotProd) { - Log.d(logTag, "Loading libsmollm_v8_4_fp16_dotprod_i8mm_sve.so") - System.loadLibrary("smollm_v8_4_fp16_dotprod_i8mm_sve") - } else if (isAtLeastArmV84 && hasSve && hasFp16 && hasDotProd) { - Log.d(logTag, "Loading libsmollm_v8_4_fp16_dotprod_sve.so") - System.loadLibrary("smollm_v8_4_fp16_dotprod_sve") - } else if (isAtLeastArmV84 && hasI8mm && hasFp16 && hasDotProd) { - Log.d(logTag, "Loading libsmollm_v8_4_fp16_dotprod_i8mm.so") - System.loadLibrary("smollm_v8_4_fp16_dotprod_i8mm") - } else if (isAtLeastArmV84 && hasFp16 && hasDotProd) { - Log.d(logTag, "Loading libsmollm_v8_4_fp16_dotprod.so") - System.loadLibrary("smollm_v8_4_fp16_dotprod") - } else if (isAtLeastArmV82 && hasFp16 && hasDotProd) { - Log.d(logTag, "Loading libsmollm_v8_2_fp16_dotprod.so") - System.loadLibrary("smollm_v8_2_fp16_dotprod") - } else if (isAtLeastArmV82 && hasFp16) { - Log.d(logTag, "Loading libsmollm_v8_2_fp16.so") - System.loadLibrary("smollm_v8_2_fp16") - } else { - Log.d(logTag, "Loading libsmollm_v8.so") - System.loadLibrary("smollm_v8") - } - } else if (Build.SUPPORTED_32_BIT_ABIS[0]?.equals("armeabi-v7a") == true) { - // armv7a (32bit) device - Log.d(logTag, "Loading libsmollm_v7a.so") - System.loadLibrary("smollm_v7a") - } else { - Log.d(logTag, "Loading default libsmollm.so") - System.loadLibrary("smollm") - } - } else { - // load the default native library with no ARM - // specific instructions - Log.d(logTag, "Loading default libsmollm.so") - System.loadLibrary("smollm") - } + // Unified single library target with Arm KleidiAI dynamic runtime micro-kernel dispatching + System.loadLibrary("smollm") } - - /** - * Reads the /proc/cpuinfo file and returns the line starting with 'Features :' that - * containing the available CPU features - */ - private fun getCPUFeatures(): String { - val cpuInfo = - try { - File("/proc/cpuinfo").readText() - } catch (e: FileNotFoundException) { - "" - } - val cpuFeatures = - cpuInfo.substringAfter("Features").substringAfter(":").substringBefore("\n").trim() - return cpuFeatures - } - - private fun supportsArm64V8a(): Boolean = Build.SUPPORTED_ABIS[0].equals("arm64-v8a") } private var nativePtr = 0L From 2c996a76544412dbd9d8666b643f8cf2075245d0 Mon Sep 17 00:00:00 2001 From: Ale Date: Wed, 16 Sep 2026 16:24:24 +0200 Subject: [PATCH 2/2] Feat: Implement direct LLM model downloading and benchmark improvements - Refactor DownloadModelsViewModel to use internal Ketch DownloadService. - Models now download directly to app's internal filesDir, preventing duplicate storage. - Automate GGUF metadata extraction and Room DB model registration upon download completion. - Add TTFT (Time To First Token) metric and cache warmup pass to C++ bench_model. --- .../model_download/DownloadModelActivity.kt | 4 +- .../model_download/DownloadModelsViewModel.kt | 80 ++++++++++++++----- smollm/src/main/cpp/LLMRunner.cpp | 18 ++++- 3 files changed, 76 insertions(+), 26 deletions(-) diff --git a/app/src/main/java/io/shubham0204/smollmandroid/ui/screens/model_download/DownloadModelActivity.kt b/app/src/main/java/io/shubham0204/smollmandroid/ui/screens/model_download/DownloadModelActivity.kt index 8decc0f9..bff734c2 100644 --- a/app/src/main/java/io/shubham0204/smollmandroid/ui/screens/model_download/DownloadModelActivity.kt +++ b/app/src/main/java/io/shubham0204/smollmandroid/ui/screens/model_download/DownloadModelActivity.kt @@ -105,7 +105,7 @@ class DownloadModelActivity : ComponentActivity() { route.modelInfo, route.modelFiles, onDownloadModel = { modelUrl -> - viewModel.downloadModelFromUrl(modelUrl) + viewModel.downloadModelFromUrl(modelUrl, onComplete = { openChatActivity() }) }, onBackClicked = { navController.navigateUp() }, ) @@ -200,7 +200,7 @@ class DownloadModelActivity : ComponentActivity() { addNewModelStep = AddNewModelStep.ImportModel }, onDownloadModelClick = { selectedPopularModelIndex -> - viewModel.downloadModelFromIndex(selectedPopularModelIndex) + viewModel.downloadModelFromIndex(selectedPopularModelIndex, onComplete = { openChatActivity() }) }, modifier = Modifier .fillMaxSize() diff --git a/app/src/main/java/io/shubham0204/smollmandroid/ui/screens/model_download/DownloadModelsViewModel.kt b/app/src/main/java/io/shubham0204/smollmandroid/ui/screens/model_download/DownloadModelsViewModel.kt index 7144f3c7..2bf04074 100644 --- a/app/src/main/java/io/shubham0204/smollmandroid/ui/screens/model_download/DownloadModelsViewModel.kt +++ b/app/src/main/java/io/shubham0204/smollmandroid/ui/screens/model_download/DownloadModelsViewModel.kt @@ -51,39 +51,75 @@ import java.net.HttpURLConnection import java.net.URL import java.nio.file.Paths +import io.shubham0204.smollmandroid.ui.screens.manage_asr.DownloadService +import kotlinx.coroutines.flow.MutableStateFlow +import kotlinx.coroutines.flow.asStateFlow +import kotlinx.coroutines.flow.update + @Single class DownloadModelsViewModel( val context: Context, val appDB: AppDB, val hfModelsAPI: HFModelsAPI, + val downloadService: DownloadService, ) : ViewModel() { - private val downloadManager = - context.getSystemService(Context.DOWNLOAD_SERVICE) as DownloadManager - fun downloadModelFromIndex(selectedPopularModelIndex: Int) { - // Downloading files in Android with the DownloadManager API - // Ref: https://youtu.be/4t8EevQSYK4?feature=shared + private val _downloadProgress = MutableStateFlow(null) + val downloadProgress = _downloadProgress.asStateFlow() + + fun downloadModelFromIndex(selectedPopularModelIndex: Int, onComplete: () -> Unit) { val modelUrl = getPopularModel(selectedPopularModelIndex)!!.url - downloadModelFromUrl(modelUrl) + downloadModelFromUrl(modelUrl, onComplete) } - fun downloadModelFromUrl(modelUrl: String) { + fun downloadModelFromUrl(modelUrl: String, onComplete: () -> Unit) { val fileName = modelUrl.substring(modelUrl.lastIndexOf('/') + 1) - val request = - DownloadManager.Request(modelUrl.toUri()) - .setTitle(fileName) - .setDescription( - "The GGUF model will be downloaded on your device for use with SmolChat." - ) - .setMimeType("application/octet-stream") - .setAllowedNetworkTypes( - DownloadManager.Request.NETWORK_WIFI or DownloadManager.Request.NETWORK_MOBILE - ) - .setNotificationVisibility( - DownloadManager.Request.VISIBILITY_VISIBLE_NOTIFY_COMPLETED - ) - .setDestinationInExternalPublicDir(Environment.DIRECTORY_DOWNLOADS, fileName) - downloadManager.enqueue(request) + val destDir = context.filesDir.absolutePath + + downloadService.startDownload( + url = modelUrl, + destDir = destDir, + destFileName = fileName, + onStart = { + Toast.makeText(context, "Starting direct download...", Toast.LENGTH_SHORT).show() + setProgressDialogTitle("Downloading Model") + setProgressDialogText("Connecting...") + showProgressDialog() + _downloadProgress.update { 0 } + }, + onProgress = { progress -> + _downloadProgress.update { progress } + setProgressDialogText("Downloading: $progress%") + }, + onSuccess = { + _downloadProgress.update { null } + setProgressDialogTitle("Registering Model") + setProgressDialogText("Analyzing GGUF metadata...") + CoroutineScope(Dispatchers.IO).launch { + val ggufReader = GGUFReader() + ggufReader.load(File(destDir, fileName).absolutePath) + val contextSize = ggufReader.getContextSize() ?: SmolLM.DefaultInferenceParams.contextSize + val chatTemplate = ggufReader.getChatTemplate() ?: SmolLM.DefaultInferenceParams.chatTemplate + appDB.addModel( + fileName, + "", + Paths.get(destDir, fileName).toString(), + contextSize.toInt(), + chatTemplate, + ) + withContext(Dispatchers.Main) { + hideProgressDialog() + Toast.makeText(context, "Model ready!", Toast.LENGTH_SHORT).show() + onComplete() + } + } + }, + onFailure = { error -> + _downloadProgress.update { null } + hideProgressDialog() + Toast.makeText(context, "Download failed: $error", Toast.LENGTH_LONG).show() + } + ) } fun getModels(query: String): Flow> = diff --git a/smollm/src/main/cpp/LLMRunner.cpp b/smollm/src/main/cpp/LLMRunner.cpp index 26622d7b..b4038cf2 100644 --- a/smollm/src/main/cpp/LLMRunner.cpp +++ b/smollm/src/main/cpp/LLMRunner.cpp @@ -332,14 +332,24 @@ std::string LLMRunner::bench_model(int pp, int tg, int pl, int nr) { const uint32_t n_ctx = llama_n_ctx(m_ctx); LOGi("bench_model: n_ctx = %u", n_ctx); + // WARMUP PASS: prime the CPU caches and KleidiAI kernels + LOGi("bench_model: running warmup pass"); + common_batch_clear(g_batch); + for (int i = 0; i < pp; i++) { + common_batch_add(g_batch, 1, i, {0}, false); + } + g_batch.logits[g_batch.n_tokens - 1] = true; + llama_decode(m_ctx, g_batch); + llama_memory_clear(llama_get_memory(m_ctx), false); + + // BENCHMARK LOOP for (int nri = 0; nri < nr; nri++) { common_batch_clear(g_batch); for (int i = 0; i < pp; i++) { common_batch_add(g_batch, 1, i, {0}, false); } g_batch.logits[g_batch.n_tokens - 1] = true; - llama_memory_clear(llama_get_memory(m_ctx), false); - + const auto t_pp_start = ggml_time_us(); llama_decode(m_ctx, g_batch); const auto t_pp_end = ggml_time_us(); @@ -379,6 +389,9 @@ std::string LLMRunner::bench_model(int pp, int tg, int pl, int nr) { pp_std = 0; tg_std = 0; } + + // TTFT is basically the time it takes to do 1 prompt processing pass (in milliseconds) + double ttft_ms = (double(pp) / pp_avg) * 1000.0; char model_desc[128]; llama_model_desc(m_model, model_desc, sizeof(model_desc)); @@ -408,6 +421,7 @@ std::string LLMRunner::bench_model(int pp, int tg, int pl, int nr) { << pp << " | " << pp_avg << " ± " << pp_std << " |\n"; result << "| " << model_desc << " | " << model_size << "GiB | " << model_n_params << "B | " << str.str() << " | tg " << tg << " | " << tg_avg << " ± " << tg_std << " |\n"; + result << "\n**TTFT (Time To First Token)**: " << ttft_ms << " ms\n"; return result.str(); }