File size: 5,076 Bytes
f2878d0 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 | // Host driver for the device engine (scalar kernel, same code as the firmware). Reads requests on
// stdin, one per line group, and prints the engine's outputs; host/conformance.mjs drives it.
// td_host <model.bin> <vocab.bin>
//
// Strings are hex-encoded UTF-8 ("-" for empty), so any text survives the line protocol.
// T <hex> -> T <n> <ids...>
// R <nq> <hex state>, then nq lines:
// Q <type 0-3> <hex text> <n_opts> <hex opt>... <has_protos 0|1|2> [lam cnt[n] center[128] vec[n*128]]
// (2: protos without a center: the older recipe)
// -> E <status> | I <T> <S> <truncated> <ms>, D <ids...>,
// and one A line per answer
#include "tinydecide.h"
#include <stdio.h>
#include <stdlib.h>
#include <string.h>
#include <string>
#include <vector>
namespace td { int debugLastIds(const uint16_t** ids); }
static std::vector<uint8_t> slurp(const char* path) {
FILE* f = fopen(path, "rb");
if (!f) { perror(path); exit(2); }
std::vector<uint8_t> b;
uint8_t buf[65536];
size_t n;
while ((n = fread(buf, 1, sizeof(buf), f)) > 0) b.insert(b.end(), buf, buf + n);
fclose(f);
return b;
}
static std::string unhex(const char* h) {
std::string s;
if (!strcmp(h, "-")) return s;
for (size_t i = 0; h[i] && h[i + 1]; i += 2) {
unsigned v;
sscanf(h + i, "%2x", &v);
s.push_back((char)v);
}
return s;
}
static char tok[1 << 20];
static const char* word() { return scanf("%1048575s", tok) == 1 ? tok : nullptr; }
int main(int argc, char** argv) {
if (argc < 3) { fprintf(stderr, "usage: td_host model.bin vocab.bin\n"); return 2; }
std::vector<uint8_t> model = slurp(argv[1]), vocab = slurp(argv[2]);
uint8_t* m = (uint8_t*)aligned_alloc(16, (model.size() + 15) & ~(size_t)15);
memcpy(m, model.data(), model.size());
if (!td::init(m, model.size(), vocab.data(), vocab.size())) { fprintf(stderr, "init failed\n"); return 2; }
static td::Answer out[td::MAX_QUESTIONS];
for (const char* w; (w = word());) {
if (!strcmp(w, "T")) {
std::string s = unhex(word());
std::vector<uint16_t> ids(s.size() + 8);
int n = td::tokenize(s.data(), s.size(), ids.data(), (int)ids.size());
printf("T %d", n);
for (int i = 0; i < n; i++) printf(" %d", ids[i]);
printf("\n");
} else if (!strcmp(w, "R")) {
const int nq = atoi(word());
std::string state = unhex(word());
std::vector<std::string> texts(nq);
std::vector<std::vector<std::string>> opts(nq);
std::vector<std::vector<const char*>> optp(nq);
std::vector<std::vector<float>> vec(nq), center(nq);
std::vector<std::vector<int>> cnt(nq);
std::vector<td::Protos> protos(nq);
std::vector<td::Question> qs(nq);
for (int k = 0; k < nq; k++) {
word(); // "Q"
const int type = atoi(word());
texts[k] = unhex(word());
const int n = atoi(word());
for (int i = 0; i < n; i++) opts[k].push_back(unhex(word()));
for (auto& o : opts[k]) optp[k].push_back(o.c_str());
const int hasP = atoi(word());
qs[k].type = (td::Type)type;
qs[k].text = texts[k].c_str();
qs[k].options = n ? optp[k].data() : nullptr;
qs[k].n_options = n;
if (hasP) {
const int K = type == td::NOUL ? 2 : n;
const float lam = (float)atof(word());
for (int i = 0; i < K; i++) cnt[k].push_back(atoi(word()));
if (hasP == 1) for (int i = 0; i < td::QDIM; i++) center[k].push_back((float)atof(word()));
for (int i = 0; i < K * td::QDIM; i++) vec[k].push_back((float)atof(word()));
protos[k] = {vec[k].data(), cnt[k].data(), hasP == 1 ? center[k].data() : nullptr, lam};
qs[k].protos = &protos[k];
}
}
td::Info info{};
const td::Status st = td::answer(state.data(), state.size(), qs.data(), nq, out, &info);
if (st != td::OK) { printf("E %d\n", (int)st); fflush(stdout); continue; }
const uint16_t* ids;
const int T = td::debugLastIds(&ids);
printf("I %d %d %d %u\nD", info.tokens, info.state_tokens, info.truncated ? 1 : 0, info.ms);
for (int i = 0; i < T; i++) printf(" %d", ids[i]);
printf("\n");
for (int k = 0; k < nq; k++) {
const td::Answer& a = out[k];
if (a.type == td::CHOICE || a.type == td::SCORE) {
printf("A c %d %d %.7g %.7g", a.n, a.pick, a.confidence, a.score);
for (int i = 0; i < a.n; i++) printf(" %.7g", a.probs[i]);
for (int i = 0; i < a.n; i++) printf(" %.7g", a.z0[i]);
} else if (a.type == td::NOUL) {
printf("A n %.7g %.7g", a.p, a.z0[1]);
} else {
printf("A s %.7g %.7g %d %d %d %d", a.p_present, a.p_span, a.tok[0], a.tok[1], a.start, a.end);
}
if (a.type != td::SPAN) for (int i = 0; i < td::QDIM; i++) printf(" %.6g", a.qvec[i]);
printf("\n");
}
}
fflush(stdout);
}
return 0;
}
|