From 1f4f19b665db7b13606154966c216eb0a4179a91 Mon Sep 17 00:00:00 2001 From: Victor Zhang Date: Wed, 10 Jun 2026 11:18:50 -0700 Subject: [PATCH] Add zstd dict trainer with ZDICT_optimizeTrainFromBuffer_cover (#806) (#806) Summary: Add a dict training framework with an abstract DictTrainer base and a concrete ZstdDictTrainer implementation. DictTrainer interface (base_dict_trainer.h): - DictNodeInfo struct: nodeID, codecID, nodeName for discovered dict-requiring nodes. - DictTrainer abstract class with findDictNodes() and trainDict() virtual methods. - trainDictsForCandidate(): top-level orchestrator that discovers dict nodes across all registered trainers, collects per-node input samples via collectInputStreams, trains dicts, packs them into the ZL_Dict wire format, patches per-node dict_id and dict_bundle_id into the serialized compressor CBOR, and returns a TrainedCandidate. trainDictsForCandidate implementation (base_dict_trainer.cpp): - Creates all registered DictTrainer instances (currently just ZstdDictTrainer). - For each trainer, calls findDictNodes() to discover nodes, then collectInputStreams() to capture input data flowing into those nodes. - For each node with samples, calls trainDict() to get packed dict content, then wraps it in Dict_pack(). - Patches the serialized compressor CBOR with per-node dict_id values and sets the bundle ID. ZstdDictTrainer (zstd_dict_trainer.h/cpp): - findDictNodes(): discovers nodes named "zl.trainable.zstd" via ZL_Compressor_forEachNode. - trainDict(): flattens MultiInput samples into byte spans, calls trainRawDict(), then packs the raw dict into the ZL_TrainedZstdContent envelope with the compression level from ZL_LocalParams. - trainRawDict(): uses ZDICT_optimizeTrainFromBuffer_cover with fallback to ZDICT_trainFromBuffer. test_dict_training.cpp: - ZstdDictTrainer unit tests: TrainDictReturnsPackedContent, TrainDictUsesCompressionLevelFromLocalParams, FindDictNodesReturnsEmptyForDefaultCompressor, FindDictNodesFindsMultipleZstdNodes. - TrainedCandidate helper tests: ReplaceBundleID, ReplaceDictID (single + batch + missing), PackFatBundleRoundTrip, PackFatBundleThrowsOnEmpty. BUCK changes: - New dict:base_dict_trainer library target. - Add dict dep to :trainer target and test deps. Differential Revision: D105006229 Pulled By: Victor-C-Zhang --- src/openzl/compress/cgraph.c | 15 + src/openzl/compress/cgraph.h | 16 + src/openzl/compress/cnodes.c | 26 ++ src/openzl/compress/cnodes.h | 5 + src/openzl/compress/nodemgr.c | 36 ++ src/openzl/compress/nodemgr.h | 5 + tools/training/BUCK | 1 + tools/training/CMakeLists.txt | 1 + tools/training/dict/BUCK | 29 ++ tools/training/dict/base_dict_trainer.cpp | 174 ++++++++++ tools/training/dict/base_dict_trainer.h | 72 ++++ tools/training/dict/zstd_dict_trainer.cpp | 157 +++++++++ tools/training/dict/zstd_dict_trainer.h | 47 +++ tools/training/tests/BUCK | 2 + tools/training/tests/test_dict_training.cpp | 350 ++++++++++++++++++++ tools/training/train.cpp | 1 + 16 files changed, 937 insertions(+) create mode 100644 tools/training/dict/BUCK create mode 100644 tools/training/dict/base_dict_trainer.cpp create mode 100644 tools/training/dict/base_dict_trainer.h create mode 100644 tools/training/dict/zstd_dict_trainer.cpp create mode 100644 tools/training/dict/zstd_dict_trainer.h create mode 100644 tools/training/tests/test_dict_training.cpp diff --git a/src/openzl/compress/cgraph.c b/src/openzl/compress/cgraph.c index 251cbaac0..79c2fad7c 100644 --- a/src/openzl/compress/cgraph.c +++ b/src/openzl/compress/cgraph.c @@ -859,6 +859,21 @@ ZL_Report ZL_Compressor_overrideGraphParams( return ZL_returnSuccess(); } +ZL_Report ZL_Compressor_overrideNodeParams( + ZL_Compressor* compressor, + ZL_NodeID node, + const ZL_NodeParameters* np) +{ + ZL_RESULT_DECLARE_SCOPE(size_t, compressor); + ZL_ERR_IF_NULL( + NM_getCNode(&compressor->nmgr, node), + node_invalid, + "Node must be registered in compressor"); + + ZL_ERR_IF_ERR(NM_overrideNodeParams(&compressor->nmgr, node, np)); + return ZL_returnSuccess(); +} + ZL_Report ZL_Compressor_overrideBaseGraph( ZL_Compressor* compressor, ZL_GraphID graph, diff --git a/src/openzl/compress/cgraph.h b/src/openzl/compress/cgraph.h index 278a4c6e2..3faa71b69 100644 --- a/src/openzl/compress/cgraph.h +++ b/src/openzl/compress/cgraph.h @@ -132,6 +132,22 @@ ZL_Report ZL_Compressor_overrideBaseGraph( ZL_GraphID graph, ZL_GraphID newBaseGraph); +/** + * Warning: This is part of experimental API for compressor mutation. + * + * Requires that: + * @p node is a parameterized node registered in @p compressor + * + * Replaces the parameters of @p node with @p np. + * @note: This function does not do any additional work typically associated + * with ZL_Compressor_registerParameterizedNode(), including dictionary + * unpacking. + */ +ZL_Report ZL_Compressor_overrideNodeParams( + ZL_Compressor* compressor, + ZL_NodeID node, + const ZL_NodeParameters* np); + /** * Look up a previously loaded dict by its ZL_DictID. * @param matDesc must match the materializer used when the dict was loaded. diff --git a/src/openzl/compress/cnodes.c b/src/openzl/compress/cnodes.c index 3717a76ca..1443199c7 100644 --- a/src/openzl/compress/cnodes.c +++ b/src/openzl/compress/cnodes.c @@ -439,3 +439,29 @@ void CTM_setDictIndex(CNodes_manager* ctm, CNodeID id, uint32_t index) ZL_ASSERT_LT(id.cnid, VECTOR_SIZE(ctm->cnodes)); VECTOR_AT(ctm->cnodes, id.cnid).maybeDictIndex = index; } + +ZL_Report CTM_overrideNodeParams( + CNodes_manager* ctm, + CNodeID id, + const ZL_NodeParameters* np) +{ + ZL_RESULT_DECLARE_SCOPE_REPORT(ctm->opCtx); + ZL_ASSERT_NN(ctm); + ZL_ASSERT_NN(np); + ZL_ASSERT_LT(id.cnid, VECTOR_SIZE(ctm->cnodes)); + + if (ZL_UniqueID_isValid(&np->mparam.mparamID.id)) { + ZL_ERR(GENERIC, "MParam override not supported"); + } + + CNode* const cnode = &VECTOR_AT(ctm->cnodes, id.cnid); + ZL_MIEncoderDesc* const trDesc = &cnode->transformDesc.publicDesc; + if (np->localParams) { + trDesc->localParams = *np->localParams; + ZL_ERR_IF_ERR(CTM_transferLocalParams(ctm, &trDesc->localParams)); + } + if (ZL_UniqueID_isValid(&np->dictID.id)) { + trDesc->dictID = np->dictID; + } + return ZL_returnSuccess(); +} diff --git a/src/openzl/compress/cnodes.h b/src/openzl/compress/cnodes.h index 5d4050578..63c7eadff 100644 --- a/src/openzl/compress/cnodes.h +++ b/src/openzl/compress/cnodes.h @@ -81,6 +81,11 @@ CTM_registerStandardTransform( /// the dict's position within the compressor's bundle. void CTM_setDictIndex(CNodes_manager* ctm, CNodeID id, uint32_t index); +ZL_Report CTM_overrideNodeParams( + CNodes_manager* ctm, + CNodeID id, + const ZL_NodeParameters* np); + /** * Rolls back the registration of @p id * @warning This only works when @p id was the last node registered. If local diff --git a/src/openzl/compress/nodemgr.c b/src/openzl/compress/nodemgr.c index 4ffb8afc9..0bb2598b4 100644 --- a/src/openzl/compress/nodemgr.c +++ b/src/openzl/compress/nodemgr.c @@ -153,6 +153,42 @@ NM_parameterizeNode(Nodes_manager* nmgr, const ZL_ParameterizedNodeDesc* desc) ZL_NodeID, NM_NodeID_fromCNodeID(ZL_RES_value(cnodeidResult))); } +ZL_Report NM_overrideNodeParams( + Nodes_manager* nmgr, + ZL_NodeID node, + const ZL_NodeParameters* np) +{ + ZL_RESULT_DECLARE_SCOPE_REPORT(nmgr->opCtx); + ZL_ASSERT_NN(nmgr); + ZL_ASSERT_NN(np); + + ZL_ERR_IF( + NM_isStandardNode(node), + node_invalid, + "Cannot replace standard node"); + + const CNodeID cnodeID = NM_CNodeID_fromNodeID(node); + ZL_ERR_IF_GE( + cnodeID.cnid, + CTM_nbCNodes(&nmgr->ctm), + node_invalid, + "Node must be registered"); + const CNode* const cnode = CTM_getCNode(&nmgr->ctm, cnodeID); + ZL_ASSERT_NN(cnode); + ZL_ERR_IF_EQ( + cnode->baseNodeID.nid, + ZL_NODE_ILLEGAL.nid, + node_invalid, + "Node must be parameterized"); + ZL_ERR_IF_NE(cnode->nodetype, node_internalTransform, node_invalid); + + if (np->name) { + ZL_ERR(parameter_invalid, "Cannot replace the name of a node"); + } + ZL_ERR_IF_ERR(CTM_overrideNodeParams(&nmgr->ctm, cnodeID, np)); + return ZL_returnSuccess(); +} + const CNode* NM_getCNode(const Nodes_manager* nmgr, ZL_NodeID nodeid) { if (NM_isStandardNode(nodeid)) { diff --git a/src/openzl/compress/nodemgr.h b/src/openzl/compress/nodemgr.h index 004236593..a97a02385 100644 --- a/src/openzl/compress/nodemgr.h +++ b/src/openzl/compress/nodemgr.h @@ -46,6 +46,11 @@ NM_registerStandardTransform( ZL_RESULT_OF(ZL_NodeID) NM_parameterizeNode(Nodes_manager* nmgr, const ZL_ParameterizedNodeDesc* desc); +ZL_Report NM_overrideNodeParams( + Nodes_manager* nmgr, + ZL_NodeID node, + const ZL_NodeParameters* np); + // Read Accessors /* diff --git a/tools/training/BUCK b/tools/training/BUCK index 1d600b0fb..597f2ff54 100644 --- a/tools/training/BUCK +++ b/tools/training/BUCK @@ -17,6 +17,7 @@ zs_cxxlibrary( "../ml_selector:ml_selector_trainer", "ace:automated_compressor_explorer", "clustering:clustering_graph_trainer", + "dict:base_dict_trainer", "graph_mutation:graph_mutation", ], exported_deps = [ diff --git a/tools/training/CMakeLists.txt b/tools/training/CMakeLists.txt index c93b090ad..40ca6a6d5 100644 --- a/tools/training/CMakeLists.txt +++ b/tools/training/CMakeLists.txt @@ -22,6 +22,7 @@ if (OPENZL_BUILD_TRAINING_TOOLS) set_property(TARGET tools_training PROPERTY POSITION_INDEPENDENT_CODE ON) target_include_directories(tools_training PUBLIC ${PROJECT_SOURCE_DIR}) + target_compile_definitions(tools_training PUBLIC ZDICT_STATIC_LINKING_ONLY) apply_openzl_compile_options_to_target(tools_training) target_link_libraries( tools_training diff --git a/tools/training/dict/BUCK b/tools/training/dict/BUCK new file mode 100644 index 000000000..513314413 --- /dev/null +++ b/tools/training/dict/BUCK @@ -0,0 +1,29 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. + +load("../../../defs.bzl", "zs_cxxlibrary") + +oncall("data_compression") + +zs_cxxlibrary( + name = "base_dict_trainer", + srcs = [ + "base_dict_trainer.cpp", + "zstd_dict_trainer.cpp", + ], + headers = [ + "base_dict_trainer.h", + "zstd_dict_trainer.h", + ], + deps = [ + "../..:logger", + "../../../cpp:openzl_cpp", + "../sample_collection:sample_collection", + ], + exported_deps = [ + "..:train_common", + "../utils:training_utils", + ], + external_deps = [ + "zstd", + ], +) diff --git a/tools/training/dict/base_dict_trainer.cpp b/tools/training/dict/base_dict_trainer.cpp new file mode 100644 index 000000000..732c3c017 --- /dev/null +++ b/tools/training/dict/base_dict_trainer.cpp @@ -0,0 +1,174 @@ +// Copyright (c) Meta Platforms, Inc. and affiliates. + +#include "tools/training/dict/base_dict_trainer.h" + +#include + +#include "openzl/compress/cgraph.h" +#include "openzl/dict/dict.h" +#include "openzl/dict/dict_constants.h" +#include "openzl/zl_compressor.h" +#include "openzl/zl_reflection.h" + +#include "tools/logger/Logger.h" +#include "tools/training/dict/zstd_dict_trainer.h" +#include "tools/training/sample_collection/training_sample_collector.h" + +using namespace openzl::tools::logger; + +namespace openzl::training { + +namespace { + +/// Build the list of all registered DictTrainer instances. +std::vector> createAllTrainers() +{ + std::vector> trainers; + trainers.push_back(std::make_unique()); + return trainers; +} + +} // anonymous namespace + +TrainedCandidate trainDictsForCandidate( + const std::vector& inputs, + Compressor& compressor, + const TrainParams& trainParams) +{ + (void)trainParams; + + TrainedCandidate candidate; + candidate.serializedCompressor = compressor.serialize(); + + // Ask each registered trainer to find its dict-requiring nodes. + auto trainers = createAllTrainers(); + + struct TrainerWork { + DictTrainer* trainer; + DictNodeInfo node; + }; + std::vector work; + std::vector nodeNames; + + for (auto& trainer : trainers) { + auto nodes = trainer->findDictNodes(compressor); + for (auto& node : nodes) { + nodeNames.push_back(node.nodeName); + work.push_back( + TrainerWork{ + .trainer = trainer.get(), + .node = std::move(node), + }); + } + } + + if (work.empty()) { + return candidate; + } + + Logger::log_c( + VERBOSE1, + "Dict training: found %zu nodes requiring dictionaries", + work.size()); + + // Use introspection hooks to collect the actual data flowing into + // each dict-requiring codec node. + auto cctx = refCCtxForTraining(compressor); + auto samplesPerNode = collectInputStreams(inputs, {}, nodeNames, cctx); + + for (auto& [trainer, node] : work) { + auto it = samplesPerNode.find(node.nodeName); + if (it == samplesPerNode.end() || it->second.empty()) { + Logger::log_c( + VERBOSE1, + "No samples collected for node %s, skipping", + node.nodeName.c_str()); + continue; + } + + const auto& nodeSamples = it->second; + Logger::log_c( + VERBOSE1, + "Training dict for node %u (codec %u, name %s) " + "with %zu samples", + node.nodeID.nid, + node.codecID, + node.nodeName.c_str(), + nodeSamples.size()); + + // trainDict returns packed dict content ready for Dict_pack. + ZL_LocalParams localParams = ZL_Compressor_Node_getLocalParams( + compressor.get(), node.nodeID); + auto dictContent = + trainer->trainDict(nodeSamples, compressor, localParams); + if (!dictContent.has_value()) { + Logger::log_c( + VERBOSE1, + "Trainer declined to train dict for node %s, skipping", + node.nodeName.c_str()); + continue; + } + + // Pack into the generic ZL_Dict wire format. + std::string packedDict(ZL_DICT_HEADER_SIZE + dictContent->size(), '\0'); + ZL_Report report = Dict_pack( + packedDict.data(), + packedDict.size(), + ZL_DICT_ID_NULL, + node.codecID, + false, + dictContent->data(), + dictContent->size()); + if (ZL_isError(report)) { + throw Exception("Dict_pack failed"); + } + packedDict.resize(ZL_validResult(report)); + + ZL_DictID generatedDictID = + Dict_extractID(packedDict.data(), packedDict.size()); + + ZL_NodeParameters nodeParams = { + .dictID = generatedDictID, + }; + compressor.unwrap(ZL_Compressor_overrideNodeParams( + compressor.get(), node.nodeID, &nodeParams)); + + candidate.dicts.push_back( + TrainedCandidate::DictEntry{ + .dictID = generatedDictID, + .packedDict = std::move(packedDict), + }); + + Logger::log_c( + VERBOSE1, + "Trained dict: %zu bytes content", + dictContent->size()); + } + + if (candidate.dicts.empty()) { + return candidate; + } + + std::vector allDictIDs; + allDictIDs.reserve(candidate.dicts.size()); + for (auto& dict : candidate.dicts) { + allDictIDs.push_back(dict.dictID); + } + // Compute bundleID + auto bundleID = + ZL_DictBundle_genBundleID(allDictIDs.data(), allDictIDs.size()); + + candidate.serializedCompressor = compressor.serialize(); + + // Set bundle ID on the candidate and patch dict_bundle_id in the CBOR. + candidate.replaceBundleID(bundleID); + + Logger::log_c( + VERBOSE1, + "Dict training complete: %zu dicts, CBOR patched with bundleID", + candidate.dicts.size()); + + return candidate; +} + +} // namespace openzl::training diff --git a/tools/training/dict/base_dict_trainer.h b/tools/training/dict/base_dict_trainer.h new file mode 100644 index 000000000..30b870bef --- /dev/null +++ b/tools/training/dict/base_dict_trainer.h @@ -0,0 +1,72 @@ +// Copyright (c) Meta Platforms, Inc. and affiliates. + +#pragma once + +#include +#include + +#include "openzl/cpp/Compressor.hpp" +#include "openzl/cpp/poly/Optional.hpp" +#include "openzl/zl_localParams.h" +#include "openzl/zl_opaque_types.h" +#include "tools/training/train_params.h" +#include "tools/training/trained_candidate.h" +#include "tools/training/utils/utils.h" + +namespace openzl::training { + +/// Info about a dict-requiring node discovered by the base trainer. +struct DictNodeInfo { + ZL_NodeID nodeID; + ZL_IDType codecID; + std::string nodeName; +}; + +/** + * Base dict trainer that trains dictionaries for a specific codec type. + * Subclasses provide a codec name for node filtering. + * Node discovery via recursive graph walking and standard-node promotion + * are handled by the base trainer. + */ +class DictTrainer { + public: + virtual ~DictTrainer() = default; + + /// Return a list of nodes to train dictionaries for. + virtual std::vector findDictNodes( + const Compressor& compressor) = 0; + + /// Train a dictionary from the given input samples and pack it into + /// the codec-specific dict content format (ready for Dict_pack()). + /// @param inputs Samples collected via introspection hooks for the + /// specific graph that feeds this codec node. + /// @param compressor The compressor that owns the node being trained. + /// @param localParams The local params of the node being trained + /// (e.g. compression level for zstd). + /// @returns Packed dict content bytes ready for Dict_pack(), or nullopt if + /// the trainer declines to train a dictionary for these inputs. + virtual poly::optional trainDict( + const std::vector& inputs, + const Compressor& compressor, + ZL_LocalParams localParams) = 0; +}; + +/// Train dictionaries for all dict-requiring nodes in @p compressor. +/// +/// For each registered DictTrainer, discovers dict-requiring nodes, uses +/// collectInputStreamsForGraphs to capture the actual data flowing into +/// each node, trains dicts, packs into a fat bundle, loads into the +/// compressor, and re-serializes. +/// +/// @param inputs Raw training samples. +/// @param compressor The compressor to train dicts for (modified +/// in place — fat bundle is loaded). +/// @param trainParams Training parameters. +/// @returns A TrainedCandidate with the re-serialized compressor, +/// bundleID, and per-dict entries populated. +TrainedCandidate trainDictsForCandidate( + const std::vector& inputs, + Compressor& compressor, + const TrainParams& trainParams); + +} // namespace openzl::training diff --git a/tools/training/dict/zstd_dict_trainer.cpp b/tools/training/dict/zstd_dict_trainer.cpp new file mode 100644 index 000000000..16f238da4 --- /dev/null +++ b/tools/training/dict/zstd_dict_trainer.cpp @@ -0,0 +1,157 @@ +// Copyright (c) Meta Platforms, Inc. and affiliates. + +#include "tools/training/dict/zstd_dict_trainer.h" + +#include + +#include "openzl/codecs/zstd/common_zstd.h" +#include "openzl/cpp/Exception.hpp" +#include "openzl/cpp/Input.hpp" +#include "openzl/zl_reflection.h" +#include "tools/logger/Logger.h" + +#include +#include + +using namespace openzl::tools::logger; + +namespace openzl::training { + +std::vector ZstdDictTrainer::findDictNodes( + const Compressor& compressor) +{ + // find all zstd nodes with the name `zl.trainable.zstd` + std::vector dictNodes; + openzl::unwrap(ZL_Compressor_forEachNode( + compressor.get(), + [](void* opaque, const ZL_Compressor* c, ZL_NodeID node) noexcept + -> ZL_Report { + const char* name = ZL_Compressor_Node_getName(c, node); + constexpr poly::string_view trainableZstdPrefix{ + "zl.trainable.zstd" + }; + if (name != nullptr + && poly::string_view(name).substr( + 0, trainableZstdPrefix.size()) + == trainableZstdPrefix) { + auto* nodes = + static_cast*>(opaque); + nodes->push_back( + DictNodeInfo{ + .nodeID = node, + .codecID = ZL_Compressor_Node_getCodecID( + c, node), + .nodeName = std::string(name), + }); + } + return ZL_returnSuccess(); + }, + &dictNodes)); + return dictNodes; +} + +poly::optional ZstdDictTrainer::trainDict( + const std::vector& inputs, + const Compressor& compressor, + ZL_LocalParams localParams) +{ + std::vector> samples; + samples.reserve(inputs.size()); + for (const auto& mi : inputs) { + for (const auto& input : *mi) { + const auto* ptr = static_cast(input.ptr()); + size_t const size = input.contentSize(); + if (ptr != nullptr && size > 0) { + samples.emplace_back(ptr, size); + } + } + } + + std::string rawDict = trainRawDict(samples); + + int32_t clevel = compressor.getParameter(CParam::CompressionLevel); + for (size_t i = 0; i < localParams.intParams.nbIntParams; ++i) { + if (localParams.intParams.intParams[i].paramId + == ZSTD_c_compressionLevel) { + clevel = localParams.intParams.intParams[i].paramValue; + break; + } + } + + size_t const contentSize = ZL_TrainedZstdContent_packedSize(rawDict.size()); + std::string packed(contentSize, '\0'); + size_t const written = ZL_TrainedZstdContent_pack( + packed.data(), + packed.size(), + clevel, + rawDict.data(), + rawDict.size()); + if (written == 0) { + throw Exception("ZL_TrainedZstdContent_pack failed"); + } + packed.resize(written); + return packed; +} + +std::string ZstdDictTrainer::trainRawDict( + const std::vector>& samples) +{ + if (samples.empty()) { + throw Exception("trainRawDict: no samples provided"); + } + + std::vector sampleSizes; + sampleSizes.reserve(samples.size()); + size_t totalSize = 0; + for (const auto& sample : samples) { + sampleSizes.push_back(sample.size()); + totalSize += sample.size(); + } + + std::vector concatenated; + concatenated.reserve(totalSize); + for (const auto& sample : samples) { + concatenated.insert(concatenated.end(), sample.begin(), sample.end()); + } + + size_t const effectiveDictSize = + std::min(maxDictSize_, std::max(totalSize / 100, minDictSize_)); + std::string dictBuffer(effectiveDictSize, '\0'); + + Logger::log_c( + VERBOSE1, + "Training zstd dict: %zu samples, %zu total bytes, " + "max dict %zu bytes", + samples.size(), + totalSize, + effectiveDictSize); + + // the train command mutates the cover params, so we need to create a copy + auto myCoverParams = coverParams_; + size_t const dictSize = ZDICT_optimizeTrainFromBuffer_cover( + dictBuffer.data(), + dictBuffer.size(), + concatenated.data(), + sampleSizes.data(), + static_cast(sampleSizes.size()), + &myCoverParams); + + if (ZDICT_isError(dictSize)) { + throw Exception( + std::string("ZDICT_optimizeTrainFromBuffer_cover failed: ") + + ZDICT_getErrorName(dictSize)); + } else { + dictBuffer.resize(dictSize); + } + + Logger::log_c( + VERBOSE1, + "Trained zstd dict: %zu bytes (cover params: d=%u, k=%u)", + dictBuffer.size(), + myCoverParams.d, + myCoverParams.k); + + return dictBuffer; +} + +} // namespace openzl::training diff --git a/tools/training/dict/zstd_dict_trainer.h b/tools/training/dict/zstd_dict_trainer.h new file mode 100644 index 000000000..b361b2b35 --- /dev/null +++ b/tools/training/dict/zstd_dict_trainer.h @@ -0,0 +1,47 @@ +// Copyright (c) Meta Platforms, Inc. and affiliates. + +#pragma once + +#include + +#include +#include +#include +#include + +#include "tools/training/dict/base_dict_trainer.h" + +namespace openzl::training { + +/** + * Trainer for zstd dictionaries. + */ +class ZstdDictTrainer : public DictTrainer { + public: + ZstdDictTrainer() = default; + ~ZstdDictTrainer() override = default; + + std::vector findDictNodes( + const Compressor& compressor) override; + + poly::optional trainDict( + const std::vector& inputs, + const Compressor& compressor, + ZL_LocalParams localParams) override; + + private: + /// Train a raw zstd dictionary from byte spans. + /// @returns Raw trained dictionary bytes (before content envelope packing). + std::string trainRawDict( + const std::vector>& samples); + + static constexpr size_t maxDictSize_{ 112 * 1024 }; + static constexpr size_t minDictSize_{ ZDICT_DICTSIZE_MIN }; + static constexpr ZDICT_cover_params_t coverParams_{ + .d = 0, + .steps = 256, + .nbThreads = 16, + }; +}; + +} // namespace openzl::training diff --git a/tools/training/tests/BUCK b/tools/training/tests/BUCK index ec10bfbd9..784044a43 100644 --- a/tools/training/tests/BUCK +++ b/tools/training/tests/BUCK @@ -14,6 +14,7 @@ cpp_unittest( "test_clustering.cpp", "test_clustering_benchmarks.cpp", "test_clustering_config_builder.cpp", + "test_dict_training.cpp", "test_genetic_algorithm.cpp", "test_sample_collection.cpp", "test_sample_limiter.cpp", @@ -33,6 +34,7 @@ cpp_unittest( "../clustering:clustering_graph_trainer", "../clustering:train_api", "../clustering:training_clustering", + "../dict:base_dict_trainer", "../utils:genetic_algorithm", "../utils:training_utils", ], diff --git a/tools/training/tests/test_dict_training.cpp b/tools/training/tests/test_dict_training.cpp new file mode 100644 index 000000000..bd565a4ed --- /dev/null +++ b/tools/training/tests/test_dict_training.cpp @@ -0,0 +1,350 @@ +// Copyright (c) Meta Platforms, Inc. and affiliates. + +#include + +#include +#include +#include + +#include "openzl/codecs/zl_zstd.h" +#include "openzl/codecs/zstd/common_zstd.h" +#include "openzl/cpp/Compressor.hpp" +#include "openzl/cpp/Exception.hpp" +#include "openzl/cpp/Input.hpp" +#include "openzl/dict/bundle.h" +#include "openzl/dict/dict.h" +#include "openzl/dict/dict_constants.h" +#include "openzl/zl_compressor.h" +#include "openzl/zl_localParams.h" +#include "openzl/zl_reflection.h" +#include "openzl/zl_unique_id.h" + +#include + +#include "tools/training/dict/base_dict_trainer.h" +#include "tools/training/dict/zstd_dict_trainer.h" +#include "tools/training/trained_candidate.h" +#include "tools/training/utils/utils.h" + +namespace openzl::training { +namespace { + +static std::vector generateSampleData( + size_t numSamples, + size_t sampleSize) +{ + std::vector samples; + samples.reserve(numSamples); + const std::string prefix = "HEADER_v1: common_prefix_data="; + const std::string suffix = " END_OF_RECORD\n"; + for (size_t i = 0; i < numSamples; ++i) { + std::string sample; + sample.reserve(sampleSize); + while (sample.size() < sampleSize) { + sample += prefix; + for (size_t j = 0; j < 32; ++j) { + sample += static_cast('A' + ((i + j) % 26)); + } + sample += suffix; + } + sample.resize(sampleSize); + samples.push_back(std::move(sample)); + } + return samples; +} + +static std::vector toMultiInputs( + const std::vector& data) +{ + std::vector inputs; + inputs.reserve(data.size()); + for (const auto& s : data) { + MultiInput mi; + mi.add(Input::refSerial(s.data(), s.size())); + inputs.push_back(std::move(mi)); + } + return inputs; +} + +// ------------------------------------------------------- +// Unit tests for ZstdDictTrainer class +// ------------------------------------------------------- + +TEST(ZstdDictTrainer, TrainDictReturnsPackedContent) +{ + auto sampleData = generateSampleData(100, 4096); + auto inputs = toMultiInputs(sampleData); + + ZstdDictTrainer trainer; + Compressor compressor; + compressor.setParameter(CParam::CompressionLevel, 3); + ZL_LocalParams localParams{}; + auto dictOpt = trainer.trainDict(inputs, compressor, localParams); + ASSERT_TRUE(dictOpt.has_value()); + std::string dictContent = dictOpt.value(); + + ASSERT_FALSE(dictContent.empty()); + ASSERT_GE(dictContent.size(), ZL_TRAINED_ZSTD_CONTENT_HEADER_SIZE); + ZL_TrainedZstdContentParsed parsed{}; + ASSERT_TRUE(ZL_TrainedZstdContent_parse( + dictContent.data(), dictContent.size(), &parsed)); + EXPECT_EQ(parsed.clevel, 3); + EXPECT_GT(parsed.rawDictSize, 0u); +} + +TEST(ZstdDictTrainer, TrainDictUsesCompressionLevelFromLocalParams) +{ + auto sampleData = generateSampleData(100, 4096); + auto inputs = toMultiInputs(sampleData); + + ZstdDictTrainer trainer; + Compressor compressor; + ZL_IntParam clevelParam = { ZSTD_c_compressionLevel, 7 }; + ZL_LocalParams localParams{ + .intParams = { .intParams = &clevelParam, .nbIntParams = 1 }, + }; + auto dictOpt = trainer.trainDict(inputs, compressor, localParams); + ASSERT_TRUE(dictOpt.has_value()); + std::string dictContent = dictOpt.value(); + + ASSERT_FALSE(dictContent.empty()); + ZL_TrainedZstdContentParsed parsed{}; + ASSERT_TRUE(ZL_TrainedZstdContent_parse( + dictContent.data(), dictContent.size(), &parsed)); + EXPECT_EQ(parsed.clevel, 7); +} + +TEST(ZstdDictTrainer, FindDictNodesReturnsEmptyForDefaultCompressor) +{ + // A default compressor has no starting graph, so the recursive + // walk finds nothing. + Compressor compressor; + auto sampleData = generateSampleData(10, 256); + auto inputs = toMultiInputs(sampleData); + TrainParams params{}; + + auto candidate = trainDictsForCandidate(inputs, compressor, params); + EXPECT_TRUE(candidate.dicts.empty()); +} + +TEST(ZstdDictTrainer, FindDictNodesFindsMultipleZstdNodes) +{ + Compressor compressor; + + // Register two distinct trainable zstd graphs, and a custom level graph + ZL_GraphID g1 = openzl::unwrap( + ZL_Compressor_buildTrainableZstdGraph(compressor.get())); + ZL_GraphID g2 = openzl::unwrap( + ZL_Compressor_buildTrainableZstdGraph(compressor.get())); + ZL_GraphID g3 = + ZL_Compressor_registerZstdGraph_withLevel(compressor.get(), 12); + ASSERT_TRUE(ZL_GraphID_isValid(g1)); + ASSERT_TRUE(ZL_GraphID_isValid(g2)); + ASSERT_TRUE(ZL_GraphID_isValid(g3)); + + // Create a split graph that sends data to these different graphs + const size_t segmentSizes[4] = { 1024, 1024, 1024, 0 }; + const ZL_GraphID successors[4] = { g1, g2, g3, ZL_GRAPH_ZSTD }; + ZL_GraphID gHead = ZL_Compressor_registerSplitGraph( + compressor.get(), ZL_Type_serial, segmentSizes, successors, 4); + ZL_Report r = ZL_Compressor_selectStartingGraphID(compressor.get(), gHead); + ASSERT_FALSE(ZL_isError(r)); + + ZstdDictTrainer zstdTrainer; + const auto nodesToTrain = zstdTrainer.findDictNodes(compressor); + ASSERT_EQ(nodesToTrain.size(), 2); + EXPECT_EQ(nodesToTrain[0].nodeName, "zl.trainable.zstd#0"); + EXPECT_EQ(nodesToTrain[1].nodeName, "zl.trainable.zstd#1"); +} + +// ------------------------------------------------------- +// Unit tests for TrainedCandidate integration helpers +// ------------------------------------------------------- + +TEST(TrainedCandidateHelpers, ReplaceBundleID) +{ + TrainedCandidate candidate; + EXPECT_FALSE(ZL_UniqueID_isValid(&candidate.bundleID.id)); + + ZL_BundleID newID; + newID.id = ZL_UniqueID_computeSHA256("test_bundle", 11); + candidate.replaceBundleID(newID); + + EXPECT_TRUE(ZL_UniqueID_eq(&candidate.bundleID.id, &newID.id)); +} + +TEST(TrainedCandidateHelpers, ReplaceDictIDUpdatesEntryAndPackedBlob) +{ + // Build a packed dict blob using Dict_pack. + const std::string content = "fake_dict_content_for_testing"; + std::string packed(ZL_DICT_HEADER_SIZE + content.size(), '\0'); + ZL_Report report = Dict_pack( + packed.data(), + packed.size(), + ZL_DICT_ID_NULL, + 42, + false, + content.data(), + content.size()); + ASSERT_FALSE(ZL_isError(report)); + packed.resize(ZL_validResult(report)); + + ZL_DictID originalID = Dict_extractID(packed.data(), packed.size()); + ASSERT_TRUE(ZL_UniqueID_isValid(&originalID.id)); + + TrainedCandidate candidate; + candidate.dicts.push_back( + TrainedCandidate::DictEntry{ + .dictID = originalID, + .packedDict = packed, + }); + + // Replace the dictID (batch of 1). + ZL_DictID newID; + newID.id = ZL_UniqueID_computeSHA256("replacement", 11); + ASSERT_FALSE(ZL_UniqueID_eq(&originalID.id, &newID.id)); + + candidate.replaceDictID({ originalID }, { newID }); + + // The struct field should be updated. + EXPECT_TRUE(ZL_UniqueID_eq(&candidate.dicts[0].dictID.id, &newID.id)); + + // The packed blob bytes 4-35 should be rewritten. + ZL_DictID extractedID = Dict_extractID( + candidate.dicts[0].packedDict.data(), + candidate.dicts[0].packedDict.size()); + EXPECT_TRUE(ZL_UniqueID_eq(&extractedID.id, &newID.id)); +} + +TEST(TrainedCandidateHelpers, ReplaceDictIDThrowsOnMiss) +{ + TrainedCandidate candidate; + ZL_DictID bogusID; + bogusID.id = ZL_UniqueID_computeSHA256("bogus", 5); + ZL_DictID anotherID; + anotherID.id = ZL_UniqueID_computeSHA256("another", 7); + + EXPECT_THROW( + candidate.replaceDictID({ bogusID }, { anotherID }), Exception); +} + +TEST(TrainedCandidateHelpers, ReplaceDictIDBatch) +{ + auto makePacked = [](const std::string& content, ZL_IDType codec) { + std::string packed(ZL_DICT_HEADER_SIZE + content.size(), '\0'); + ZL_Report report = Dict_pack( + packed.data(), + packed.size(), + ZL_DICT_ID_NULL, + codec, + false, + content.data(), + content.size()); + EXPECT_FALSE(ZL_isError(report)); + packed.resize(ZL_validResult(report)); + return packed; + }; + + std::string packed1 = makePacked("dict_alpha_AAAA", 42); + std::string packed2 = makePacked("dict_bravo_BBBB", 43); + ZL_DictID id1 = Dict_extractID(packed1.data(), packed1.size()); + ZL_DictID id2 = Dict_extractID(packed2.data(), packed2.size()); + + TrainedCandidate candidate; + candidate.dicts.push_back( + TrainedCandidate::DictEntry{ + .dictID = id1, + .packedDict = std::move(packed1), + }); + candidate.dicts.push_back( + TrainedCandidate::DictEntry{ + .dictID = id2, + .packedDict = std::move(packed2), + }); + + ZL_DictID new1, new2; + new1.id = ZL_UniqueID_computeSHA256("new1", 4); + new2.id = ZL_UniqueID_computeSHA256("new2", 4); + + candidate.replaceDictID({ id1, id2 }, { new1, new2 }); + + EXPECT_TRUE(ZL_UniqueID_eq(&candidate.dicts[0].dictID.id, &new1.id)); + EXPECT_TRUE(ZL_UniqueID_eq(&candidate.dicts[1].dictID.id, &new2.id)); + + ZL_DictID ext1 = Dict_extractID( + candidate.dicts[0].packedDict.data(), + candidate.dicts[0].packedDict.size()); + ZL_DictID ext2 = Dict_extractID( + candidate.dicts[1].packedDict.data(), + candidate.dicts[1].packedDict.size()); + EXPECT_TRUE(ZL_UniqueID_eq(&ext1.id, &new1.id)); + EXPECT_TRUE(ZL_UniqueID_eq(&ext2.id, &new2.id)); +} + +TEST(TrainedCandidateHelpers, PackFatBundleRoundTrip) +{ + // Create two packed dicts. + const std::string content1 = "dict_content_one_AAAAAAAAA"; + const std::string content2 = "dict_content_two_BBBBBBBBB"; + + auto makePacked = [](const std::string& content, ZL_IDType codec) { + std::string packed(ZL_DICT_HEADER_SIZE + content.size(), '\0'); + ZL_Report report = Dict_pack( + packed.data(), + packed.size(), + ZL_DICT_ID_NULL, + codec, + false, + content.data(), + content.size()); + EXPECT_FALSE(ZL_isError(report)); + packed.resize(ZL_validResult(report)); + return packed; + }; + + std::string packed1 = makePacked(content1, 42); + std::string packed2 = makePacked(content2, 43); + + ZL_DictID id1 = Dict_extractID(packed1.data(), packed1.size()); + ZL_DictID id2 = Dict_extractID(packed2.data(), packed2.size()); + + TrainedCandidate candidate; + candidate.dicts.push_back( + TrainedCandidate::DictEntry{ + .dictID = id1, + .packedDict = std::move(packed1), + }); + candidate.dicts.push_back( + TrainedCandidate::DictEntry{ + .dictID = id2, + .packedDict = std::move(packed2), + }); + candidate.bundleID = ZL_DictBundle_genBundleID( + std::vector{ id1, id2 }.data(), 2); + + // Pack the fat bundle. + std::string fatBundle = candidate.packFatBundle(); + ASSERT_FALSE(fatBundle.empty()); + + // Parse the BundleInfo header from the fat bundle. + auto parseResult = ZL_BundleInfo_parse(fatBundle.data(), fatBundle.size()); + ASSERT_FALSE(ZL_RES_isError(parseResult)); + + ZL_BundleInfo info = ZL_RES_value(parseResult); + EXPECT_TRUE(info.isFatBundle); + EXPECT_EQ(info.numDicts, 2u); + EXPECT_TRUE(ZL_UniqueID_eq(&info.bundleID.id, &candidate.bundleID.id)); + + // The dict IDs in the header should match. + EXPECT_TRUE(ZL_UniqueID_eq(&info.dictIDs[0].id, &id1.id)); + EXPECT_TRUE(ZL_UniqueID_eq(&info.dictIDs[1].id, &id2.id)); +} + +TEST(TrainedCandidateHelpers, PackFatBundleThrowsOnEmpty) +{ + TrainedCandidate candidate; + EXPECT_THROW(candidate.packFatBundle(), Exception); +} + +} // namespace +} // namespace openzl::training diff --git a/tools/training/train.cpp b/tools/training/train.cpp index 6dcdcf21a..479eb03ea 100644 --- a/tools/training/train.cpp +++ b/tools/training/train.cpp @@ -7,6 +7,7 @@ #include "tools/ml_selector/ml_selector_trainer.h" #include "tools/training/ace/ace.h" #include "tools/training/clustering/clustering_graph_trainer.h" +#include "tools/training/dict/base_dict_trainer.h" #include "tools/training/graph_mutation/graph_mutation_utils.h" #include "tools/training/train.h" #include "tools/training/utils/serialized_compressor_internal.h"