// Set-BPE serialization: quantized cell trajectories <-> token streams. // // A port of the runtime's SpatialVocabulary / TemporalVocabulary (Python, // ``tokenization/setbpe_{spatial,temporal}.py``). Every merge is applied from // a priority queue ordered exactly as the Python heaps order their tuples, so // the token streams are identical, not merely equivalent. #include #include #include #include #include #include #include #include #include #include #include #include #include #include "parallel.hpp" namespace py = pybind11; namespace { int lowest_bit(uint64_t mask) { return mask ? __builtin_ctzll(mask) : -1; } uint64_t pair_key(int64_t left, int64_t right) { return (static_cast(left) << 32) | static_cast(right); } struct Rule { int64_t child; int64_t rank; }; class SetBPE { public: SetBPE(std::vector bins, int64_t lookahead, const std::vector>& spatial_merges, const std::vector>& temporal_merges) : bins_(std::move(bins)), lookahead_(lookahead) { const int64_t dimension = static_cast(bins_.size()); if (dimension < 1 || dimension > 63) throw std::invalid_argument("bins must have 1..63 slots"); if (lookahead_ < 1) throw std::invalid_argument("lookahead must be positive"); offsets_.assign(dimension + 1, 0); for (int64_t d = 0; d < dimension; ++d) { if (bins_[d] < 1) throw std::invalid_argument("every dimension needs at least one bin"); offsets_[d + 1] = offsets_[d] + bins_[d]; } atom_count_ = offsets_[dimension]; spatial_size_ = atom_count_ + static_cast(spatial_merges.size()); support_.assign(spatial_size_, 0); expansion_.assign(spatial_size_ * dimension, -1); for (int64_t d = 0; d < dimension; ++d) { for (int64_t v = 0; v < bins_[d]; ++v) { const int64_t token = offsets_[d] + v; support_[token] = uint64_t{1} << d; expansion_[token * dimension + d] = v; } } for (size_t rank = 0; rank < spatial_merges.size(); ++rank) { const auto [left, right] = spatial_merges[rank]; const int64_t token = atom_count_ + static_cast(rank); if (left < 0 || right < 0 || left >= token || right >= token) throw std::invalid_argument("a merge rule must reference already-created tokens"); if (support_[left] & support_[right]) throw std::invalid_argument("a merge rule joins two overlapping supports"); support_[token] = support_[left] | support_[right]; for (int64_t d = 0; d < dimension; ++d) { expansion_[token * dimension + d] = std::max(expansion_[left * dimension + d], expansion_[right * dimension + d]); } spatial_rules_[pair_key(std::min(left, right), std::max(left, right))] = Rule{token, static_cast(rank)}; } spatial_anchor_.resize(spatial_size_); for (int64_t token = 0; token < spatial_size_; ++token) spatial_anchor_[token] = lowest_bit(support_[token]); // Rooted temporal tokens: a pattern of (time offset, spatial token) pairs. vocab_size_ = spatial_size_ + static_cast(temporal_merges.size()); patterns_.resize(vocab_size_); root_.resize(vocab_size_); for (int64_t token = 0; token < spatial_size_; ++token) { patterns_[token] = {{0, token}}; root_[token] = token; } for (size_t rank = 0; rank < temporal_merges.size(); ++rank) { const auto [left, right, offset] = temporal_merges[rank]; const int64_t child = spatial_size_ + static_cast(rank); if (left < 0 || right < 0 || left >= child || right >= child) throw std::invalid_argument("a temporal merge must reference already-created tokens"); if (offset < 1 || offset >= lookahead_) throw std::invalid_argument("a temporal merge offset must be in [1, lookahead)"); std::vector> pattern = patterns_[left]; for (const auto& [time, token] : patterns_[right]) pattern.emplace_back(time + offset, token); std::sort(pattern.begin(), pattern.end()); const int64_t span = 1 + pattern.back().first; if (span > lookahead_) throw std::invalid_argument("temporal merge exceeds the lookahead horizon"); std::vector masks(span, 0); int roots = 0; for (const auto& [time, token] : pattern) { if (masks[time] & support_[token]) throw std::invalid_argument("temporal merge joins overlapping spatial supports"); masks[time] |= support_[token]; roots += time == 0; } if (roots != 1) throw std::invalid_argument("a rooted temporal token must have exactly one time-zero spatial token"); patterns_[child] = std::move(pattern); root_[child] = root_[left]; const auto key = std::make_tuple(left, right, offset); if (temporal_rules_.count(key)) throw std::invalid_argument("duplicate temporal merge rule"); temporal_rules_[key] = Rule{child, static_cast(rank)}; } } int64_t vocab_size() const { return vocab_size_; } // One dense [T, D] trajectory with -1 for semantically absent slots. std::vector encode(const int64_t* cells, int64_t steps) const { const int64_t dimension = static_cast(bins_.size()); if (steps < 1) throw std::invalid_argument("every episode must have at least one step"); uint64_t first_support = 0; for (int64_t t = 0; t < steps; ++t) { uint64_t support = 0; for (int64_t d = 0; d < dimension; ++d) { const int64_t value = cells[t * dimension + d]; if (value < -1 || value >= bins_[d]) throw std::invalid_argument("episodes must contain -1 or an in-range cell index"); if (value >= 0) support |= uint64_t{1} << d; } if (!support) throw std::invalid_argument("every timestep needs at least one active coordinate"); if (t == 0) first_support = support; else if (support != first_support) throw std::invalid_argument("active dimensions must remain constant within an episode"); } // Instances, as in the Python encoder: token, anchor step, liveness. std::vector tokens, anchors; std::vector active; std::vector>> at(steps); // (token, instance) for (int64_t t = 0; t < steps; ++t) { for (int64_t token : spatial_encode(cells + t * dimension)) { at[t].emplace_back(token, static_cast(tokens.size())); tokens.push_back(token); anchors.push_back(t); active.push_back(1); } } // (rank, anchor, root anchor dimension, left instance, right instance, child) using Offer = std::tuple; std::priority_queue, std::greater> frontier; auto offer = [&](int64_t left_instance, int64_t right_instance) { const auto found = temporal_rules_.find(std::make_tuple( tokens[left_instance], tokens[right_instance], anchors[right_instance] - anchors[left_instance])); if (found == temporal_rules_.end()) return; const Rule& rule = found->second; frontier.emplace(rule.rank, anchors[left_instance], spatial_anchor_[root_[rule.child]], left_instance, right_instance, rule.child); }; for (int64_t anchor = 0; anchor < steps; ++anchor) { for (int64_t offset = 1; offset < std::min(lookahead_, steps - anchor); ++offset) { for (const auto& [left_token, left_instance] : at[anchor]) { for (const auto& [right_token, right_instance] : at[anchor + offset]) { offer(left_instance, right_instance); } } } } auto erase = [&](int64_t step, int64_t token) { auto& items = at[step]; for (auto it = items.begin(); it != items.end(); ++it) { if (it->first == token) { items.erase(it); return; } } }; while (!frontier.empty()) { const auto [rank, anchor, root_anchor, left_instance, right_instance, child] = frontier.top(); frontier.pop(); (void)rank; (void)root_anchor; if (!active[left_instance] || !active[right_instance]) continue; if (anchors[left_instance] != anchor) continue; const auto resolved = temporal_rules_.find(std::make_tuple( tokens[left_instance], tokens[right_instance], anchors[right_instance] - anchor)); if (resolved == temporal_rules_.end() || resolved->second.child != child) continue; active[left_instance] = 0; active[right_instance] = 0; erase(anchor, tokens[left_instance]); erase(anchors[right_instance], tokens[right_instance]); const int64_t new_instance = static_cast(tokens.size()); tokens.push_back(child); anchors.push_back(anchor); active.push_back(1); at[anchor].emplace_back(child, new_instance); for (int64_t offset = 1; offset < lookahead_; ++offset) { const int64_t future = anchor + offset; if (future < steps) { for (const auto& [token, other] : at[future]) offer(new_instance, other); } const int64_t past = anchor - offset; if (past >= 0) { for (const auto& [token, other] : at[past]) offer(other, new_instance); } } } std::vector> ordered; for (int64_t anchor = 0; anchor < steps; ++anchor) { for (const auto& [token, instance] : at[anchor]) { if (active[instance]) ordered.emplace_back(anchor, spatial_anchor_[root_[token]], token); } } std::sort(ordered.begin(), ordered.end()); std::vector stream; stream.reserve(ordered.size()); for (const auto& item : ordered) stream.push_back(std::get<2>(item)); return stream; } // Tokens -> dense [T, D] cells (-1 outside ``full``), validating like the runtime. std::vector decode(const int64_t* stream, int64_t length, uint64_t full) const { const int64_t dimension = static_cast(bins_.size()); std::vector coverage{0}; std::vector> spatial_tokens(1); size_t front = 0; for (int64_t index = 0; index < length; ++index) { const int64_t token = stream[index]; if (token < 0 || token >= vocab_size_) throw std::invalid_argument("temporal token " + std::to_string(token) + " is outside the vocabulary"); while (front < coverage.size() && coverage[front] == full) ++front; if (front == coverage.size()) { coverage.push_back(0); spatial_tokens.emplace_back(); } const uint64_t uncovered = full & ~coverage[front]; if (spatial_anchor_[root_[token]] != lowest_bit(uncovered)) throw std::invalid_argument("temporal stream violates root order at the earliest incomplete step"); for (const auto& [offset, spatial_token] : patterns_[token]) { const size_t step = front + static_cast(offset); while (step >= coverage.size()) { coverage.push_back(0); spatial_tokens.emplace_back(); } if (coverage[step] & support_[spatial_token]) throw std::invalid_argument("temporal stream contains overlapping spatial support"); coverage[step] |= support_[spatial_token]; spatial_tokens[step].push_back(spatial_token); } } for (uint64_t mask : coverage) { if (mask != full) throw std::invalid_argument("temporal stream ended with an incomplete control step"); } std::vector cells(coverage.size() * dimension, -1); for (size_t step = 0; step < coverage.size(); ++step) { for (int64_t token : spatial_tokens[step]) { for (int64_t d = 0; d < dimension; ++d) { const int64_t value = expansion_[token * dimension + d]; if (value >= 0) cells[step * dimension + d] = value; } } } return cells; } private: // Spatial Set-BPE of one dense row: merges in rank order, then anchor order. std::vector spatial_encode(const int64_t* row) const { const int64_t dimension = static_cast(bins_.size()); std::vector present; for (int64_t d = 0; d < dimension; ++d) { if (row[d] >= 0) present.push_back(offsets_[d] + row[d]); } std::sort(present.begin(), present.end()); using Candidate = std::tuple; // rank, child, left, right std::priority_queue, std::greater> heap; for (size_t i = 0; i < present.size(); ++i) { for (size_t j = i + 1; j < present.size(); ++j) { const auto found = spatial_rules_.find(pair_key(present[i], present[j])); if (found != spatial_rules_.end()) heap.emplace(found->second.rank, found->second.child, present[i], present[j]); } } std::vector tokens = present; // a small set; linear lookups are fastest auto contains = [&](int64_t token) { return std::find(tokens.begin(), tokens.end(), token) != tokens.end(); }; while (!heap.empty()) { const auto [rank, child, left, right] = heap.top(); heap.pop(); (void)rank; if (!contains(left) || !contains(right)) continue; tokens.erase(std::find(tokens.begin(), tokens.end(), left)); tokens.erase(std::find(tokens.begin(), tokens.end(), right)); for (int64_t other : tokens) { const int64_t low = std::min(child, other), high = std::max(child, other); const auto found = spatial_rules_.find(pair_key(low, high)); if (found != spatial_rules_.end()) heap.emplace(found->second.rank, found->second.child, low, high); } tokens.push_back(child); } std::sort(tokens.begin(), tokens.end(), [&](int64_t a, int64_t b) { return std::make_pair(spatial_anchor_[a], a) < std::make_pair(spatial_anchor_[b], b); }); return tokens; } struct TupleHash { size_t operator()(const std::tuple& key) const { const auto [a, b, c] = key; return std::hash()((static_cast(a) * 1000003u) ^ (static_cast(b) << 20) ^ static_cast(c)); } }; std::vector bins_; int64_t lookahead_; std::vector offsets_; int64_t atom_count_ = 0, spatial_size_ = 0, vocab_size_ = 0; std::vector support_; std::vector expansion_; std::vector spatial_anchor_; std::unordered_map spatial_rules_; std::vector>> patterns_; std::vector root_; std::unordered_map, Rule, TupleHash> temporal_rules_; }; // Encode B dense [T, D] trajectories in parallel; returns one stream per row. py::list encode_batch(const SetBPE& codec, const py::array_t& cells, int threads) { if (cells.ndim() != 3) throw std::invalid_argument("cells must have shape [B, T, D]"); const int64_t batch = cells.shape(0), steps = cells.shape(1), width = cells.shape(2); const int64_t* data = cells.data(); std::vector> streams(batch); { py::gil_scoped_release release; ac2::parallel_for(batch, threads, [&](int64_t index) { streams[index] = codec.encode(data + index * steps * width, steps); }); } py::list output(batch); for (int64_t index = 0; index < batch; ++index) output[index] = py::cast(streams[index]); return output; } // Decode B equal-horizon streams to dense [B, T, D] cells (-1 outside ``full``). py::array_t decode_batch(const SetBPE& codec, const std::vector>& streams, int64_t width, uint64_t full, int threads) { const int64_t batch = static_cast(streams.size()); if (batch < 1) throw std::invalid_argument("streams must not be empty"); std::vector> cells(batch); { py::gil_scoped_release release; ac2::parallel_for(batch, threads, [&](int64_t index) { if (streams[index].empty()) throw std::invalid_argument("token stream must not be empty"); cells[index] = codec.decode(streams[index].data(), static_cast(streams[index].size()), full); }); } const int64_t steps = static_cast(cells[0].size()) / width; for (const auto& row : cells) { if (static_cast(row.size()) != steps * width) throw std::invalid_argument("all token streams in a batch must decode to the same length"); } py::array_t output({batch, steps, width}); int64_t* out = output.mutable_data(); for (int64_t index = 0; index < batch; ++index) std::copy(cells[index].begin(), cells[index].end(), out + index * steps * width); return output; } } // namespace void register_setbpe(py::module_& module) { py::class_(module, "SetBPE") .def(py::init, int64_t, const std::vector>&, const std::vector>&>(), py::arg("bins"), py::arg("lookahead"), py::arg("spatial_merges"), py::arg("temporal_merges")) .def_property_readonly("vocab_size", &SetBPE::vocab_size) .def("encode_batch", &encode_batch, py::arg("cells"), py::arg("threads") = 0, "Dense [B, T, D] cells (-1 = absent) -> one token stream per row.") .def("decode_batch", &decode_batch, py::arg("streams"), py::arg("width"), py::arg("full_support"), py::arg("threads") = 0, "Equal-horizon streams -> dense [B, T, D] cells, -1 outside full_support."); }