Ling-3.0-tiny-RKNN / src /service.cpp
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
22.4 kB
#include "ling3/service.h"
#include "ling3/decoder.h"
#include "ling3/model_package.h"
#include "ling3/tokenizer.h"
#include <algorithm>
#include <atomic>
#include <cctype>
#include <charconv>
#include <chrono>
#include <cmath>
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <cstdlib>
#include <functional>
#include <iomanip>
#include <iostream>
#include <limits>
#include <mutex>
#include <sstream>
#include <span>
#include <stdexcept>
#include <string>
#include <string_view>
#include <thread>
#include <utility>
#include <vector>
#if defined(__linux__)
#include <arpa/inet.h>
#include <cerrno>
#include <csignal>
#include <netinet/in.h>
#include <sys/socket.h>
#include <unistd.h>
#endif
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;
};
#if defined(__linux__)
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);
}
}
#endif
} // namespace
int RunHttpService(
const std::filesystem::path & package_path,
std::uint16_t port,
std::string_view bind_address) {
#if !defined(__linux__)
(void)package_path;
(void)port;
(void)bind_address;
throw std::runtime_error("the HTTP service currently requires Linux");
#else
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();
}
#endif
}
} // namespace ling3