Download esp32/host/td_host.cpp from TheREZOR/TinyDecide: direct link, hf CLI and curl.
- Browser
- Download file 5.08 kB
-
https://huggingface.co/TheREZOR/TinyDecide/resolve/main/esp32/host/td_host.cpp
- Command line
-
hf download hf://TheREZOR/TinyDecide/esp32/host/td_host.cpp
-
curl -L -o td_host.cpp https://huggingface.co/TheREZOR/TinyDecide/resolve/main/esp32/host/td_host.cpp
5.08 kB
| // 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 | |
| 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; | |
| } | |