Ling-3.0-tiny-RKNN / tools /gdn_prefill_compare.cpp
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
12 kB
#include <rknn_api.h>
#include <algorithm>
#include <chrono>
#include <cmath>
#include <cstdint>
#include <cstdio>
#include <cstring>
#include <fstream>
#include <limits>
#include <stdexcept>
#include <string>
#include <string_view>
#include <vector>
namespace {
void Check(int status, const char * operation) {
if (status != RKNN_SUCC) {
throw std::runtime_error(
std::string(operation) + " failed with status " + std::to_string(status));
}
}
std::vector<std::uint8_t> ReadBytes(const std::string & path) {
std::ifstream stream(path, std::ios::binary | std::ios::ate);
if (!stream) throw std::runtime_error("cannot open " + path);
const auto size = stream.tellg();
std::vector<std::uint8_t> result(static_cast<std::size_t>(size));
stream.seekg(0);
stream.read(reinterpret_cast<char *>(result.data()), size);
if (!stream) throw std::runtime_error("cannot read " + path);
return result;
}
std::vector<float> ReadFloats(const std::string & path) {
const auto bytes = ReadBytes(path);
if (bytes.size() % sizeof(float) != 0) {
throw std::runtime_error("invalid float data " + path);
}
std::vector<float> result(bytes.size() / sizeof(float));
std::memcpy(result.data(), bytes.data(), bytes.size());
return result;
}
std::size_t Elements(const rknn_tensor_attr & value) {
if (value.n_elems != 0) return value.n_elems;
std::size_t count = 1;
for (std::uint32_t index = 0; index < value.n_dims; ++index) count *= value.dims[index];
return count;
}
struct Context {
rknn_context value = 0;
std::vector<rknn_tensor_attr> inputs;
std::vector<rknn_tensor_attr> outputs;
std::vector<rknn_tensor_mem *> input_memory;
std::vector<rknn_tensor_mem *> output_memory;
Context(const std::string & path, int core) {
const auto model = ReadBytes(path);
Check(rknn_init(&value, const_cast<std::uint8_t *>(model.data()), model.size(), 0, nullptr),
"initialize GDN context");
Check(rknn_set_core_mask(value, static_cast<rknn_core_mask>(1U << core)),
"set GDN core");
rknn_input_output_num counts {};
Check(rknn_query(value, RKNN_QUERY_IN_OUT_NUM, &counts, sizeof(counts)), "query GDN counts");
inputs.resize(counts.n_input);
outputs.resize(counts.n_output);
input_memory.resize(counts.n_input);
output_memory.resize(counts.n_output);
for (std::uint32_t index = 0; index < counts.n_input; ++index) {
inputs[index].index = index;
Check(rknn_query(value, RKNN_QUERY_INPUT_ATTR, &inputs[index], sizeof(inputs[index])),
"query GDN input");
input_memory[index] = rknn_create_mem2(
value, std::max(inputs[index].size, inputs[index].size_with_stride),
RKNN_FLAG_MEMORY_CACHEABLE);
if (input_memory[index] == nullptr) throw std::runtime_error("allocate GDN input");
auto binding = inputs[index];
binding.pass_through = 1;
Check(rknn_set_io_mem(value, input_memory[index], &binding), "bind GDN input");
}
for (std::uint32_t index = 0; index < counts.n_output; ++index) {
outputs[index].index = index;
Check(rknn_query(value, RKNN_QUERY_OUTPUT_ATTR, &outputs[index], sizeof(outputs[index])),
"query GDN output");
output_memory[index] = rknn_create_mem2(
value, std::max(outputs[index].size, outputs[index].size_with_stride),
RKNN_FLAG_MEMORY_CACHEABLE);
if (output_memory[index] == nullptr) throw std::runtime_error("allocate GDN output");
auto binding = outputs[index];
binding.pass_through = 1;
Check(rknn_set_io_mem(value, output_memory[index], &binding), "bind GDN output");
}
}
Context(const Context &) = delete;
Context & operator=(const Context &) = delete;
~Context() {
if (value == 0) return;
for (auto * memory : output_memory) if (memory != nullptr) rknn_destroy_mem(value, memory);
for (auto * memory : input_memory) if (memory != nullptr) rknn_destroy_mem(value, memory);
rknn_destroy(value);
}
int Input(std::string_view name) const {
for (std::size_t index = 0; index < inputs.size(); ++index) {
if (name == inputs[index].name) return static_cast<int>(index);
}
throw std::runtime_error("missing GDN input " + std::string(name));
}
int Output(std::string_view name) const {
for (std::size_t index = 0; index < outputs.size(); ++index) {
if (name == outputs[index].name) return static_cast<int>(index);
}
throw std::runtime_error("missing GDN output " + std::string(name));
}
void Stage(int index, const float * input, std::size_t count) {
if (Elements(inputs[index]) != count) {
throw std::runtime_error("GDN input shape or type mismatch");
}
const auto & attribute = inputs[index];
if (attribute.type == RKNN_TENSOR_FLOAT16) {
auto * output = static_cast<__fp16 *>(input_memory[index]->virt_addr);
for (std::size_t i = 0; i < count; ++i) output[i] = static_cast<__fp16>(input[i]);
} else if (attribute.type == RKNN_TENSOR_INT8) {
auto * output = static_cast<std::int8_t *>(input_memory[index]->virt_addr);
for (std::size_t i = 0; i < count; ++i) {
const auto quantized = static_cast<int>(
std::lround(input[i] / attribute.scale)) + attribute.zp;
output[i] = static_cast<std::int8_t>(std::clamp(quantized, -128, 127));
}
} else {
throw std::runtime_error("unsupported GDN input dtype");
}
Check(rknn_mem_sync(value, input_memory[index], RKNN_MEMORY_SYNC_TO_DEVICE),
"sync GDN input");
}
void Run() { Check(rknn_run(value, nullptr), "run GDN"); }
std::vector<float> Read(int index) {
Check(rknn_mem_sync(value, output_memory[index], RKNN_MEMORY_SYNC_FROM_DEVICE),
"sync GDN output");
const auto count = Elements(outputs[index]);
std::vector<float> result(count);
const auto & attribute = outputs[index];
if (attribute.type == RKNN_TENSOR_FLOAT16) {
const auto * input = static_cast<const __fp16 *>(output_memory[index]->virt_addr);
for (std::size_t i = 0; i < count; ++i) result[i] = static_cast<float>(input[i]);
} else if (attribute.type == RKNN_TENSOR_INT8) {
const auto * input = static_cast<const std::int8_t *>(output_memory[index]->virt_addr);
for (std::size_t i = 0; i < count; ++i) {
result[i] = (static_cast<int>(input[i]) - attribute.zp) * attribute.scale;
}
} else {
throw std::runtime_error("unsupported GDN output dtype");
}
return result;
}
};
struct Error {
double cosine = 0.0;
double mean = 0.0;
double maximum = 0.0;
};
Error Compare(const std::vector<float> & actual, const std::vector<float> & expected) {
if (actual.size() != expected.size()) throw std::runtime_error("comparison size mismatch");
double dot = 0.0, actual_norm = 0.0, expected_norm = 0.0, mean = 0.0, maximum = 0.0;
for (std::size_t index = 0; index < actual.size(); ++index) {
const double a = actual[index];
const double b = expected[index];
const double delta = std::abs(a - b);
dot += a * b;
actual_norm += a * a;
expected_norm += b * b;
mean += delta;
maximum = std::max(maximum, delta);
}
return {dot / std::sqrt(actual_norm * expected_norm), mean / actual.size(), maximum};
}
} // namespace
int main(int argc, char ** argv) try {
if (argc != 7 && argc != 8) {
std::fprintf(
stderr,
"usage: %s SINGLE.rknn BATCH.rknn DATA_DIR HEADS CORE TOKENS [ITERATIONS]\n",
argv[0]);
return 2;
}
constexpr int dimension = 128;
const int heads = std::atoi(argv[4]);
const int core = std::atoi(argv[5]);
const int tokens = std::atoi(argv[6]);
const int iterations = argc == 8 ? std::atoi(argv[7]) : 20;
if (heads < 1 || core < 0 || core > 2 || tokens < 1 || iterations < 1) {
throw std::invalid_argument("heads, core, tokens, or iterations is invalid");
}
const std::size_t width = static_cast<std::size_t>(heads) * dimension;
const std::size_t state_elements = width * dimension;
const std::string data = argv[3];
const auto query = ReadFloats(data + "/query.f32");
const auto key = ReadFloats(data + "/key.f32");
const auto value = ReadFloats(data + "/value.f32");
const auto decay = ReadFloats(data + "/decay.f32");
const auto beta = ReadFloats(data + "/beta.f32");
auto state = ReadFloats(data + "/state.f32");
Context single(argv[1], core);
std::vector<float> sequential_output(static_cast<std::size_t>(tokens) * width);
for (int token = 0; token < tokens; ++token) {
const auto offset = static_cast<std::size_t>(token) * width;
single.Stage(single.Input("query"), query.data() + offset, width);
single.Stage(single.Input("key"), key.data() + offset, width);
single.Stage(single.Input("value"), value.data() + offset, width);
single.Stage(single.Input("decay"), decay.data() + offset, width);
single.Stage(single.Input("beta"), beta.data() + static_cast<std::size_t>(token) * heads, heads);
single.Stage(single.Input("state"), state.data(), state_elements);
single.Run();
const auto output = single.Read(single.Output("output"));
std::copy(output.begin(), output.end(), sequential_output.begin() + offset);
state = single.Read(single.Output("new_state"));
}
Context batch(argv[2], core);
batch.Stage(batch.Input("query"), query.data(), query.size());
batch.Stage(batch.Input("key"), key.data(), key.size());
batch.Stage(batch.Input("value"), value.data(), value.size());
batch.Stage(batch.Input("decay"), decay.data(), decay.size());
batch.Stage(batch.Input("beta"), beta.data(), beta.size());
const auto initial_state = ReadFloats(data + "/state.f32");
batch.Stage(batch.Input("state"), initial_state.data(), initial_state.size());
batch.Run();
const auto batch_output = batch.Read(batch.Output("output"));
const auto batch_state = batch.Read(batch.Output("new_state"));
const auto output_error = Compare(batch_output, ReadFloats(data + "/output.f32"));
const auto state_error = Compare(batch_state, ReadFloats(data + "/new_state.f32"));
batch.Run();
const auto benchmark_begin = std::chrono::steady_clock::now();
for (int iteration = 0; iteration < iterations; ++iteration) batch.Run();
const auto benchmark_end = std::chrono::steady_clock::now();
const double batch_run_ms = std::chrono::duration<double, std::milli>(
benchmark_end - benchmark_begin).count() / iterations;
std::printf(
"OUTPUT cosine=%.9f mean_abs=%.9g max_abs=%.9g\n"
"STATE cosine=%.9f mean_abs=%.9g max_abs=%.9g\n"
"BATCH tokens=%d iterations=%d rknn_run_ms=%.6f tokens_per_second=%.3f"
" input_dtype=%s output_dtype=%s\n",
output_error.cosine, output_error.mean, output_error.maximum,
state_error.cosine, state_error.mean, state_error.maximum,
tokens, iterations, batch_run_ms, 1000.0 * tokens / batch_run_ms,
batch.inputs.front().type == RKNN_TENSOR_INT8 ? "int8" : "fp16",
batch.outputs.front().type == RKNN_TENSOR_INT8 ? "int8" : "fp16");
return output_error.cosine >= 0.999 && state_error.cosine >= 0.999 ? 0 : 3;
} catch (const std::exception & error) {
std::fprintf(stderr, "error: %s\n", error.what());
return 1;
}