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;
}