// 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 // // Strings are hex-encoded UTF-8 ("-" for empty), so any text survives the line protocol. // T -> T // R , then nq lines: // Q ... [lam cnt[n] center[128] vec[n*128]] // (2: protos without a center: the older recipe) // -> E | I , D , // and one A line per answer #include "tinydecide.h" #include #include #include #include #include namespace td { int debugLastIds(const uint16_t** ids); } static std::vector slurp(const char* path) { FILE* f = fopen(path, "rb"); if (!f) { perror(path); exit(2); } std::vector 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 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 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 texts(nq); std::vector> opts(nq); std::vector> optp(nq); std::vector> vec(nq), center(nq); std::vector> cnt(nq); std::vector protos(nq); std::vector 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; }