Download src/service.cpp from Sariel00/Ling-3.0-tiny-RKNN: direct link, hf CLI and curl.
- Browser
- Download file 22.4 kB
-
https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/src/service.cpp
- Command line
-
hf download hf://Sariel00/Ling-3.0-tiny-RKNN/src/service.cpp
-
curl -L -o service.cpp https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/src/service.cpp
22.4 kB
| namespace ling3 { | |
| namespace { | |
| using Clock = std::chrono::steady_clock; | |
| double Milliseconds(Clock::time_point begin, Clock::time_point end) { | |
| return std::chrono::duration<double, std::milli>(end - begin).count(); | |
| } | |
| template <typename T> | |
| std::span<const T> TensorSpan(const TensorView & tensor) { | |
| if (tensor.entry->data_bytes % sizeof(T) != 0) { | |
| throw std::runtime_error(std::string(tensor.name) + " has an incompatible byte count"); | |
| } | |
| return { | |
| reinterpret_cast<const T *>(tensor.data), | |
| static_cast<std::size_t>(tensor.entry->data_bytes / sizeof(T)), | |
| }; | |
| } | |
| Tokenizer PackageTokenizer(const ModelPackage & package) { | |
| const auto & asset = package.tensor("tokenizer"); | |
| if (asset.entry->role != static_cast<std::uint32_t>(TensorRole::kTokenizer)) { | |
| throw std::runtime_error("tokenizer tensor has the wrong role"); | |
| } | |
| return Tokenizer(TensorSpan<std::byte>(asset)); | |
| } | |
| std::string JsonEscape(std::string_view input) { | |
| std::string output; | |
| output.reserve(input.size() + 8); | |
| constexpr char hex[] = "0123456789abcdef"; | |
| for (const unsigned char value : input) { | |
| switch (value) { | |
| case '"': output += "\\\""; break; | |
| case '\\': output += "\\\\"; break; | |
| case '\b': output += "\\b"; break; | |
| case '\f': output += "\\f"; break; | |
| case '\n': output += "\\n"; break; | |
| case '\r': output += "\\r"; break; | |
| case '\t': output += "\\t"; break; | |
| default: | |
| if (value < 0x20) { | |
| output += "\\u00"; | |
| output.push_back(hex[value >> 4U]); | |
| output.push_back(hex[value & 0x0FU]); | |
| } else { | |
| output.push_back(static_cast<char>(value)); | |
| } | |
| } | |
| } | |
| return output; | |
| } | |
| std::string JsonNumber(double value) { | |
| std::ostringstream stream; | |
| stream << std::fixed << std::setprecision(3) << value; | |
| return stream.str(); | |
| } | |
| class SessionEngine { | |
| public: | |
| SessionEngine( | |
| const ModelPackage & package, | |
| const Tokenizer & tokenizer, | |
| Decoder & decoder, | |
| double initialization_ms) | |
| : package_(package), | |
| tokenizer_(tokenizer), | |
| decoder_(decoder), | |
| logits_(package.header().vocab_size), | |
| initialization_ms_(initialization_ms) {} | |
| bool active() const noexcept { return active_.load(std::memory_order_relaxed); } | |
| std::size_t position() const noexcept { | |
| return position_.load(std::memory_order_relaxed); | |
| } | |
| std::size_t turns() const noexcept { return turns_.load(std::memory_order_relaxed); } | |
| double initialization_ms() const noexcept { return initialization_ms_; } | |
| bool Cancel() noexcept { | |
| if (!active()) return false; | |
| cancel_requested_.store(true, std::memory_order_relaxed); | |
| return true; | |
| } | |
| void Reset() { | |
| cancel_requested_.store(true, std::memory_order_relaxed); | |
| std::lock_guard lock(inference_mutex_); | |
| decoder_.Reset(); | |
| position_.store(0, std::memory_order_relaxed); | |
| turns_.store(0, std::memory_order_relaxed); | |
| cancel_requested_.store(false, std::memory_order_relaxed); | |
| } | |
| void Chat( | |
| std::string_view user_text, | |
| std::size_t maximum_new_tokens, | |
| const std::function<bool(std::string_view)> & emit) { | |
| const auto request_begin = Clock::now(); | |
| std::unique_lock lock(inference_mutex_); | |
| const auto lock_acquired = Clock::now(); | |
| cancel_requested_.store(false, std::memory_order_relaxed); | |
| active_.store(true, std::memory_order_relaxed); | |
| struct ActiveGuard { | |
| std::atomic_bool & active; | |
| ~ActiveGuard() { active.store(false, std::memory_order_relaxed); } | |
| } guard {active_}; | |
| if (user_text.empty()) throw std::invalid_argument("request body must contain user text"); | |
| if (maximum_new_tokens == 0 || maximum_new_tokens > 2048) { | |
| throw std::invalid_argument("max_tokens must be in [1, 2048]"); | |
| } | |
| const bool first_turn = decoder_.position() == 0; | |
| const std::string prompt = first_turn | |
| ? "<role>SYSTEM</role>detailed thinking off<|role_end|>" | |
| "<role>HUMAN</role>" + std::string(user_text) + | |
| "<|role_end|><role>ASSISTANT</role>\n<think></think>" | |
| : "<|role_end|><role>HUMAN</role>" + std::string(user_text) + | |
| "<|role_end|><role>ASSISTANT</role>\n<think></think>"; | |
| const auto prompt_tokens = tokenizer_.Encode(prompt); | |
| if (prompt_tokens.empty() || | |
| decoder_.position() + prompt_tokens.size() + maximum_new_tokens > | |
| package_.header().max_context) { | |
| throw std::runtime_error("request would exceed the package context capacity"); | |
| } | |
| const std::string start_event = | |
| "{\"type\":\"start\",\"turn\":" + std::to_string(turns() + 1) + | |
| ",\"position\":" + std::to_string(decoder_.position()) + | |
| ",\"prompt_tokens\":" + std::to_string(prompt_tokens.size()) + | |
| ",\"queue_ms\":" + JsonNumber(Milliseconds(request_begin, lock_acquired)) + "}\n"; | |
| if (!emit(start_event)) return; | |
| const auto prefill_begin = Clock::now(); | |
| double prefill_eval_ms = 0.0; | |
| std::size_t prefill_offset = 0; | |
| std::size_t prefill_batch_count = 0; | |
| std::size_t batch32_count = 0; | |
| std::size_t batched_tokens = 0; | |
| std::size_t maximum_batch_rows = 0; | |
| const std::size_t batch_granularity = decoder_.batch_granularity(); | |
| while (batch_granularity != 0 && prompt_tokens.size() - prefill_offset >= 2) { | |
| std::size_t rows = std::min<std::size_t>( | |
| prompt_tokens.size() - prefill_offset, 128); | |
| rows -= rows % batch_granularity; | |
| if (rows < 2) break; | |
| const auto block = std::span<const std::uint32_t>(prompt_tokens).subspan( | |
| prefill_offset, rows); | |
| const bool needs_logits = prompt_tokens.size() - prefill_offset == rows; | |
| prefill_eval_ms += (needs_logits | |
| ? decoder_.EvalBatch(block, logits_) | |
| : decoder_.EvalBatchState(block)).total_ms; | |
| prefill_offset += rows; | |
| ++prefill_batch_count; | |
| batch32_count += rows == 32 ? 1 : 0; | |
| batched_tokens += rows; | |
| maximum_batch_rows = std::max(maximum_batch_rows, rows); | |
| position_.store(decoder_.position(), std::memory_order_relaxed); | |
| } | |
| for (; prefill_offset < prompt_tokens.size(); ++prefill_offset) { | |
| prefill_eval_ms += decoder_.Eval(prompt_tokens[prefill_offset], logits_).total_ms; | |
| position_.store(decoder_.position(), std::memory_order_relaxed); | |
| } | |
| const auto prefill_end = Clock::now(); | |
| std::size_t generated_tokens = 0; | |
| double decode_eval_ms = 0.0; | |
| bool stopped_on_eos = false; | |
| bool client_connected = true; | |
| const auto generation_begin = Clock::now(); | |
| double first_token_ms = 0.0; | |
| for (std::size_t index = 0; index < maximum_new_tokens; ++index) { | |
| if (cancel_requested_.load(std::memory_order_relaxed)) break; | |
| const auto found = std::max_element(logits_.begin(), logits_.end()); | |
| const auto token = static_cast<std::uint32_t>(found - logits_.begin()); | |
| if (token == package_.header().eos_token) { | |
| stopped_on_eos = true; | |
| break; | |
| } | |
| if (generated_tokens == 0) { | |
| first_token_ms = Milliseconds(request_begin, Clock::now()); | |
| } | |
| const std::string event = | |
| "{\"type\":\"token\",\"id\":" + std::to_string(token) + | |
| ",\"text\":\"" + JsonEscape(tokenizer_.Piece(token)) + | |
| "\",\"index\":" + std::to_string(generated_tokens) + | |
| ",\"elapsed_ms\":" + JsonNumber(Milliseconds(request_begin, Clock::now())) + "}\n"; | |
| if (!emit(event)) { | |
| client_connected = false; | |
| cancel_requested_.store(true, std::memory_order_relaxed); | |
| break; | |
| } | |
| ++generated_tokens; | |
| const auto timing = decoder_.Eval(token, logits_); | |
| decode_eval_ms += timing.total_ms; | |
| position_.store(decoder_.position(), std::memory_order_relaxed); | |
| } | |
| turns_.fetch_add(1, std::memory_order_relaxed); | |
| const bool canceled = cancel_requested_.load(std::memory_order_relaxed); | |
| const double prefill_ms = Milliseconds(prefill_begin, prefill_end); | |
| const double generation_ms = Milliseconds(generation_begin, Clock::now()); | |
| if (client_connected) { | |
| const std::string finish_event = | |
| "{\"type\":\"finish\",\"reason\":\"" + | |
| std::string(canceled ? "canceled" : (stopped_on_eos ? "eos" : "length")) + | |
| "\",\"turn\":" + std::to_string(turns()) + | |
| ",\"position\":" + std::to_string(decoder_.position()) + | |
| ",\"prompt_tokens\":" + std::to_string(prompt_tokens.size()) + | |
| ",\"prefill_batch_count\":" + std::to_string(prefill_batch_count) + | |
| ",\"prefill_batch32_count\":" + std::to_string(batch32_count) + | |
| ",\"prefill_batched_tokens\":" + std::to_string(batched_tokens) + | |
| ",\"prefill_max_batch_rows\":" + std::to_string(maximum_batch_rows) + | |
| ",\"generated_tokens\":" + std::to_string(generated_tokens) + | |
| ",\"queue_ms\":" + JsonNumber(Milliseconds(request_begin, lock_acquired)) + | |
| ",\"prefill_ms\":" + JsonNumber(prefill_ms) + | |
| ",\"prefill_eval_ms\":" + JsonNumber(prefill_eval_ms) + | |
| ",\"prefill_tokens_per_second\":" + | |
| JsonNumber(prefill_ms == 0.0 ? 0.0 : 1000.0 * prompt_tokens.size() / prefill_ms) + | |
| ",\"ttft_ms\":" + JsonNumber(first_token_ms) + | |
| ",\"generation_ms\":" + JsonNumber(generation_ms) + | |
| ",\"decode_eval_ms\":" + JsonNumber(decode_eval_ms) + | |
| ",\"decode_tokens_per_second\":" + | |
| JsonNumber(decode_eval_ms == 0.0 ? 0.0 : 1000.0 * generated_tokens / decode_eval_ms) + | |
| "}\n"; | |
| emit(finish_event); | |
| } | |
| } | |
| private: | |
| const ModelPackage & package_; | |
| const Tokenizer & tokenizer_; | |
| Decoder & decoder_; | |
| std::vector<float> logits_; | |
| std::mutex inference_mutex_; | |
| std::atomic_bool active_ {false}; | |
| std::atomic_bool cancel_requested_ {false}; | |
| std::atomic_size_t position_ {0}; | |
| std::atomic_size_t turns_ {0}; | |
| double initialization_ms_ = 0.0; | |
| }; | |
| struct HttpRequest { | |
| std::string method; | |
| std::string target; | |
| std::string body; | |
| }; | |
| class Socket { | |
| public: | |
| explicit Socket(int descriptor = -1) : descriptor_(descriptor) {} | |
| ~Socket() { if (descriptor_ >= 0) close(descriptor_); } | |
| Socket(const Socket &) = delete; | |
| Socket & operator=(const Socket &) = delete; | |
| Socket(Socket && other) noexcept : descriptor_(std::exchange(other.descriptor_, -1)) {} | |
| int get() const noexcept { return descriptor_; } | |
| private: | |
| int descriptor_; | |
| }; | |
| bool SendAll(int socket, std::string_view data) { | |
| while (!data.empty()) { | |
| const auto sent = send(socket, data.data(), data.size(), MSG_NOSIGNAL); | |
| if (sent < 0 && errno == EINTR) continue; | |
| if (sent <= 0) return false; | |
| data.remove_prefix(static_cast<std::size_t>(sent)); | |
| } | |
| return true; | |
| } | |
| void SendResponse( | |
| int socket, | |
| int status, | |
| std::string_view reason, | |
| std::string_view body, | |
| std::string_view content_type = "application/json; charset=utf-8") { | |
| const std::string header = | |
| "HTTP/1.1 " + std::to_string(status) + " " + std::string(reason) + "\r\n" | |
| "Content-Type: " + std::string(content_type) + "\r\n" | |
| "Content-Length: " + std::to_string(body.size()) + "\r\n" | |
| "Connection: close\r\n\r\n"; | |
| SendAll(socket, header); | |
| SendAll(socket, body); | |
| } | |
| bool SendChunk(int socket, std::string_view data) { | |
| std::ostringstream size; | |
| size << std::hex << data.size() << "\r\n"; | |
| return SendAll(socket, size.str()) && SendAll(socket, data) && SendAll(socket, "\r\n"); | |
| } | |
| std::size_t ParseContentLength(std::string_view headers) { | |
| std::size_t cursor = headers.find("\r\n") + 2; | |
| while (cursor < headers.size()) { | |
| const auto end = headers.find("\r\n", cursor); | |
| if (end == std::string_view::npos || end == cursor) break; | |
| const auto line = headers.substr(cursor, end - cursor); | |
| const auto colon = line.find(':'); | |
| if (colon != std::string_view::npos) { | |
| std::string name(line.substr(0, colon)); | |
| std::transform(name.begin(), name.end(), name.begin(), [](unsigned char value) { | |
| return static_cast<char>(std::tolower(value)); | |
| }); | |
| if (name == "content-length") { | |
| auto value = line.substr(colon + 1); | |
| while (!value.empty() && value.front() == ' ') value.remove_prefix(1); | |
| std::size_t length = 0; | |
| const auto [end_ptr, error] = | |
| std::from_chars(value.data(), value.data() + value.size(), length); | |
| if (error != std::errc {} || end_ptr != value.data() + value.size()) { | |
| throw std::runtime_error("invalid Content-Length"); | |
| } | |
| return length; | |
| } | |
| } | |
| cursor = end + 2; | |
| } | |
| return 0; | |
| } | |
| HttpRequest ReadRequest(int socket) { | |
| constexpr std::size_t kMaximumHeaders = 64 * 1024; | |
| constexpr std::size_t kMaximumBody = 1024 * 1024; | |
| std::string data; | |
| data.reserve(4096); | |
| std::size_t header_end = std::string::npos; | |
| char buffer[4096]; | |
| while ((header_end = data.find("\r\n\r\n")) == std::string::npos) { | |
| const auto count = recv(socket, buffer, sizeof(buffer), 0); | |
| if (count < 0 && errno == EINTR) continue; | |
| if (count <= 0) throw std::runtime_error("connection closed before request headers"); | |
| data.append(buffer, static_cast<std::size_t>(count)); | |
| if (data.size() > kMaximumHeaders) throw std::runtime_error("request headers are too large"); | |
| } | |
| header_end += 4; | |
| const std::string_view headers(data.data(), header_end); | |
| const auto request_line_end = headers.find("\r\n"); | |
| const auto first_space = headers.find(' '); | |
| const auto second_space = first_space == std::string_view::npos | |
| ? std::string_view::npos : headers.find(' ', first_space + 1); | |
| if (request_line_end == std::string_view::npos || first_space == std::string_view::npos || | |
| second_space == std::string_view::npos || second_space > request_line_end) { | |
| throw std::runtime_error("invalid HTTP request line"); | |
| } | |
| HttpRequest request; | |
| request.method = std::string(headers.substr(0, first_space)); | |
| request.target = std::string(headers.substr(first_space + 1, second_space - first_space - 1)); | |
| const auto body_bytes = ParseContentLength(headers); | |
| if (body_bytes > kMaximumBody) throw std::runtime_error("request body is too large"); | |
| while (data.size() - header_end < body_bytes) { | |
| const auto count = recv(socket, buffer, sizeof(buffer), 0); | |
| if (count < 0 && errno == EINTR) continue; | |
| if (count <= 0) throw std::runtime_error("connection closed before request body"); | |
| data.append(buffer, static_cast<std::size_t>(count)); | |
| } | |
| request.body.assign(data.data() + header_end, body_bytes); | |
| return request; | |
| } | |
| std::string_view Path(std::string_view target) { | |
| const auto query = target.find('?'); | |
| return target.substr(0, query); | |
| } | |
| std::size_t MaxTokens(std::string_view target) { | |
| constexpr std::string_view key = "max_tokens="; | |
| const auto found = target.find(key); | |
| if (found == std::string_view::npos) return 64; | |
| const auto begin = found + key.size(); | |
| const auto end = target.find('&', begin); | |
| const auto value = target.substr(begin, end == std::string_view::npos ? end : end - begin); | |
| std::size_t result = 0; | |
| const auto [end_ptr, error] = | |
| std::from_chars(value.data(), value.data() + value.size(), result); | |
| if (error != std::errc {} || end_ptr != value.data() + value.size()) { | |
| throw std::invalid_argument("invalid max_tokens query parameter"); | |
| } | |
| return result; | |
| } | |
| void HandleClient(Socket client, SessionEngine & engine) { | |
| try { | |
| const auto request = ReadRequest(client.get()); | |
| const auto path = Path(request.target); | |
| if (request.method == "GET" && path == "/health") { | |
| const std::string body = | |
| "{\"status\":\"ok\",\"active\":" + | |
| std::string(engine.active() ? "true" : "false") + | |
| ",\"position\":" + std::to_string(engine.position()) + | |
| ",\"turns\":" + std::to_string(engine.turns()) + | |
| ",\"initialization_ms\":" + JsonNumber(engine.initialization_ms()) + "}\n"; | |
| SendResponse(client.get(), 200, "OK", body); | |
| return; | |
| } | |
| if (request.method == "POST" && path == "/v1/cancel") { | |
| const bool canceled = engine.Cancel(); | |
| SendResponse(client.get(), 200, "OK", | |
| canceled ? "{\"cancel_requested\":true}\n" : | |
| "{\"cancel_requested\":false}\n"); | |
| return; | |
| } | |
| if (request.method == "POST" && path == "/v1/reset") { | |
| engine.Reset(); | |
| SendResponse(client.get(), 200, "OK", "{\"reset\":true}\n"); | |
| return; | |
| } | |
| if (request.method == "POST" && path == "/v1/chat") { | |
| const std::string header = | |
| "HTTP/1.1 200 OK\r\n" | |
| "Content-Type: application/x-ndjson; charset=utf-8\r\n" | |
| "Transfer-Encoding: chunked\r\n" | |
| "Cache-Control: no-cache\r\n" | |
| "Connection: close\r\n\r\n"; | |
| if (!SendAll(client.get(), header)) return; | |
| try { | |
| engine.Chat(request.body, MaxTokens(request.target), [&](std::string_view event) { | |
| return SendChunk(client.get(), event); | |
| }); | |
| } catch (const std::exception & error) { | |
| const std::string event = | |
| "{\"type\":\"error\",\"message\":\"" + JsonEscape(error.what()) + "\"}\n"; | |
| SendChunk(client.get(), event); | |
| } | |
| SendAll(client.get(), "0\r\n\r\n"); | |
| return; | |
| } | |
| SendResponse(client.get(), 404, "Not Found", "{\"error\":\"not found\"}\n"); | |
| } catch (const std::exception & error) { | |
| const std::string body = | |
| "{\"error\":\"" + JsonEscape(error.what()) + "\"}\n"; | |
| SendResponse(client.get(), 400, "Bad Request", body); | |
| } | |
| } | |
| } // namespace | |
| int RunHttpService( | |
| const std::filesystem::path & package_path, | |
| std::uint16_t port, | |
| std::string_view bind_address) { | |
| (void)package_path; | |
| (void)port; | |
| (void)bind_address; | |
| throw std::runtime_error("the HTTP service currently requires Linux"); | |
| std::signal(SIGPIPE, SIG_IGN); | |
| const auto initialize_begin = Clock::now(); | |
| const ModelPackage package(package_path); | |
| const auto tokenizer = PackageTokenizer(package); | |
| Decoder decoder(package); | |
| const std::size_t batch_granularity = decoder.batch_granularity(); | |
| if (batch_granularity != 0) { | |
| for (std::size_t rows = batch_granularity; rows <= 128; rows *= 2) { | |
| decoder.PrepareBatch(rows); | |
| } | |
| } | |
| const double initialization_ms = Milliseconds(initialize_begin, Clock::now()); | |
| SessionEngine engine(package, tokenizer, decoder, initialization_ms); | |
| Socket server(socket(AF_INET, SOCK_STREAM | SOCK_CLOEXEC, 0)); | |
| if (server.get() < 0) { | |
| throw std::runtime_error(std::string("socket: ") + std::strerror(errno)); | |
| } | |
| int reuse = 1; | |
| setsockopt(server.get(), SOL_SOCKET, SO_REUSEADDR, &reuse, sizeof(reuse)); | |
| sockaddr_in address {}; | |
| address.sin_family = AF_INET; | |
| address.sin_port = htons(port); | |
| const std::string bind_text(bind_address); | |
| if (inet_pton(AF_INET, bind_text.c_str(), &address.sin_addr) != 1) { | |
| throw std::invalid_argument("bind address must be an IPv4 address"); | |
| } | |
| if (bind(server.get(), reinterpret_cast<const sockaddr *>(&address), sizeof(address)) != 0) { | |
| throw std::runtime_error(std::string("bind: ") + std::strerror(errno)); | |
| } | |
| if (listen(server.get(), 16) != 0) { | |
| throw std::runtime_error(std::string("listen: ") + std::strerror(errno)); | |
| } | |
| std::cout << std::fixed << std::setprecision(3) | |
| << "service_ready=http://" << bind_address << ':' << port << '\n' | |
| << "initialization_ms=" << initialization_ms << '\n' | |
| << "prewarm_experts=" | |
| << (std::getenv("LING3_PREWARM_EXPERTS") == nullptr ? "false" : "true") | |
| << '\n' | |
| << "dynamic_batch_1_128=" << (decoder.has_dynamic_batch() ? "true" : "false") | |
| << std::endl; | |
| while (true) { | |
| const int descriptor = accept4(server.get(), nullptr, nullptr, SOCK_CLOEXEC); | |
| if (descriptor < 0 && errno == EINTR) continue; | |
| if (descriptor < 0) { | |
| throw std::runtime_error(std::string("accept: ") + std::strerror(errno)); | |
| } | |
| std::thread([descriptor, &engine]() { | |
| HandleClient(Socket(descriptor), engine); | |
| }).detach(); | |
| } | |
| } | |
| } // namespace ling3 | |