From a941b9555381b942e37d696a6095bab2340f73bb Mon Sep 17 00:00:00 2001 From: Melody Ren Date: Tue, 11 Aug 2026 19:30:31 -0700 Subject: [PATCH] Allow a decoder to be constructed from a raw Stim DEM A decoder entry may name a Stim detector error model instead of the matrix form of the same model. The decoder is then constructed from the DEM text, which is the only construction a DEM-native decoder such as Chromobius accepts, and derives its own dimensions and observable mapping from it. The matrix keys become optional at the parser and the checks that compare against syndrome_size move to the constructed decoder; the matrix path is otherwise unchanged. The model text also reaches a nested global decoder through global_decoder_params, so trt_decoder can build Chromobius rather than handing it a parity check matrix it cannot use. Signed-off-by: Melody Ren --- .../cudaq/qec/realtime/decoding_config.h | 5 + .../plugins/trt_decoder/trt_decoder.cpp | 9 +- libs/qec/lib/realtime/config.cpp | 34 +++++-- libs/qec/lib/realtime/realtime_decoding.cpp | 69 +++++++++++--- .../python/bindings/py_decoding_config.cpp | 1 + libs/qec/unittests/CMakeLists.txt | 10 ++ libs/qec/unittests/test_decoders_yaml.cpp | 95 +++++++++++++++++++ 7 files changed, 199 insertions(+), 24 deletions(-) diff --git a/libs/qec/include/cudaq/qec/realtime/decoding_config.h b/libs/qec/include/cudaq/qec/realtime/decoding_config.h index 4fd6448ea..6fd2b682d 100644 --- a/libs/qec/include/cudaq/qec/realtime/decoding_config.h +++ b/libs/qec/include/cudaq/qec/realtime/decoding_config.h @@ -72,6 +72,11 @@ struct decoder_config { std::optional cuda_device_id; uint64_t block_size = 0; uint64_t syndrome_size = 0; + /// Path to a Stim detector error model. When set, the decoder is + /// constructed from the DEM text rather than from `H_sparse`, which is what + /// a DEM-native decoder such as Chromobius requires. Interpreted like the + /// other model paths in a configuration, relative to the working directory. + std::string stim_dem_path; std::vector H_sparse; std::vector O_sparse; std::vector D_sparse; diff --git a/libs/qec/lib/decoders/plugins/trt_decoder/trt_decoder.cpp b/libs/qec/lib/decoders/plugins/trt_decoder/trt_decoder.cpp index 675b25fd9..4f0e41d17 100644 --- a/libs/qec/lib/decoders/plugins/trt_decoder/trt_decoder.cpp +++ b/libs/qec/lib/decoders/plugins/trt_decoder/trt_decoder.cpp @@ -762,8 +762,15 @@ trt_decoder::trt_decoder(const cudaq::qec::sparse_binary_matrix &H, global_decoder_params_ = params.get("global_decoder_params"); if (!global_decoder_name.empty()) { + // A DEM-native global decoder is constructed from the model text; the + // matrix arm is what every other global decoder takes. global_decoder_ = - decoder::get(global_decoder_name, H, global_decoder_params_); + global_decoder_params_.contains("stim_dem") + ? decoder::get( + global_decoder_name, + global_decoder_params_.get("stim_dem"), + global_decoder_params_) + : decoder::get(global_decoder_name, H, global_decoder_params_); CUDA_QEC_INFO("TensorRT decoder: global_decoder '{}' attached", global_decoder_name); } diff --git a/libs/qec/lib/realtime/config.cpp b/libs/qec/lib/realtime/config.cpp index 3cfaddf07..5b8684177 100644 --- a/libs/qec/lib/realtime/config.cpp +++ b/libs/qec/lib/realtime/config.cpp @@ -281,17 +281,25 @@ struct MappingTraits { io.mapOptional("dispatch", config.dispatch, cudaq::qec::decoding::config::DecoderDispatch::host); io.mapOptional("cuda_device_id", config.cuda_device_id); - io.mapRequired("block_size", config.block_size); - io.mapRequired("syndrome_size", config.syndrome_size); - io.mapRequired("H_sparse", config.H_sparse); - io.mapRequired("O_sparse", config.O_sparse); + io.mapOptional("stim_dem_path", config.stim_dem_path, std::string{}); + // A DEM-sourced decoder derives these, so they cannot be required at the + // parser; create_realtime_decoder() decides which model keys it needs. + io.mapOptional("block_size", config.block_size, std::uint64_t{0}); + io.mapOptional("syndrome_size", config.syndrome_size, std::uint64_t{0}); + io.mapOptional("H_sparse", config.H_sparse, std::vector{}); + io.mapOptional("O_sparse", config.O_sparse, std::vector{}); io.mapRequired("D_sparse", config.D_sparse); + // A DEM-sourced decoder derives its own dimensions, so the checks below + // that compare against syndrome_size do not apply to it; the D row count + // is checked against the constructed decoder instead. + const bool from_dem = !config.stim_dem_path.empty(); + // Validate that the number of rows in the H_sparse vector is equal to // syndrome_size. auto num_H_rows = std::count(config.H_sparse.begin(), config.H_sparse.end(), -1); - if (num_H_rows != config.syndrome_size) { + if (!from_dem && num_H_rows != config.syndrome_size) { throw std::runtime_error( "Number of rows in H_sparse vector is not equal to syndrome_size: " + std::to_string(num_H_rows) + @@ -320,7 +328,7 @@ struct MappingTraits { if (!config.D_sparse.empty()) { auto num_D_rows = std::count(config.D_sparse.begin(), config.D_sparse.end(), -1); - if (num_D_rows != config.syndrome_size) { + if (!from_dem && num_D_rows != config.syndrome_size) { throw std::runtime_error("Number of rows in D_sparse vector is not " "equal to syndrome_size: " + std::to_string(num_D_rows) + @@ -608,6 +616,7 @@ std::string decoder_config_json_schema() { {"block_size", llvm::json::Object{{"type", "integer"}, {"minimum", 0}}}, {"syndrome_size", llvm::json::Object{{"type", "integer"}, {"minimum", 0}}}, + {"stim_dem_path", llvm::json::Object{{"type", "string"}}}, {"H_sparse", llvm::json::Object{{"$ref", "#/$defs/sparse_matrix"}}}, {"O_sparse", llvm::json::Object{{"$ref", "#/$defs/sparse_matrix"}}}, {"D_sparse", llvm::json::Object{{"$ref", "#/$defs/sparse_matrix"}}}, @@ -683,9 +692,16 @@ std::string decoder_config_json_schema() { llvm::json::Object{ {"type", "object"}, {"properties", std::move(config_properties)}, - {"required", - llvm::json::Array{"id", "type", "block_size", "syndrome_size", - "H_sparse", "O_sparse", "D_sparse"}}, + {"required", llvm::json::Array{"id", "type", "D_sparse"}}, + // Either a raw DEM or the matrix form of the same model. + {"oneOf", + llvm::json::Array{ + llvm::json::Object{ + {"required", + llvm::json::Array{"block_size", "syndrome_size", + "H_sparse", "O_sparse"}}}, + llvm::json::Object{ + {"required", llvm::json::Array{"stim_dem_path"}}}}}, {"additionalProperties", false}, {"allOf", std::move(dispatch)}}}, {"decoder_params", std::move(decoder_params)}, diff --git a/libs/qec/lib/realtime/realtime_decoding.cpp b/libs/qec/lib/realtime/realtime_decoding.cpp index 63baa2774..c918c81bc 100644 --- a/libs/qec/lib/realtime/realtime_decoding.cpp +++ b/libs/qec/lib/realtime/realtime_decoding.cpp @@ -18,6 +18,8 @@ #include #include #include +#include +#include #include #include #include @@ -156,6 +158,17 @@ static std::vector pack_syndrome_bits(const uint8_t *syndromes, namespace cudaq::qec::decoding::host { +/// Read the DEM a decoder entry names. +std::string read_stim_dem( + const cudaq::qec::decoding::config::decoder_config &decoder_config) { + std::ifstream file(decoder_config.stim_dem_path); + if (!file) + throw std::runtime_error(fmt::format( + "stim_dem_path could not be opened: {}", decoder_config.stim_dem_path)); + return std::string((std::istreambuf_iterator(file)), + std::istreambuf_iterator()); +} + cudaqx::heterogeneous_map prepare_decoder_params( const cudaq::qec::decoding::config::decoder_config &decoder_config) { auto params = decoder_config.decoder_custom_args_to_heterogeneous_map(); @@ -192,6 +205,16 @@ cudaqx::heterogeneous_map prepare_decoder_params( if (has_global_decoder && !params.contains("global_decoder_params")) params.insert("global_decoder_params", cudaqx::heterogeneous_map()); + // A DEM-native global decoder (chromobius) cannot be built from the parent's + // H, so pass the model text down beside it. Decoders that do not ask for it + // ignore the key. + if (has_global_decoder && !decoder_config.stim_dem_path.empty()) { + auto global_decoder_params = + params.get("global_decoder_params"); + global_decoder_params.insert("stim_dem", read_stim_dem(decoder_config)); + params.insert("global_decoder_params", global_decoder_params); + } + if (decoder_config.O_sparse.empty()) return params; @@ -231,26 +254,44 @@ std::unique_ptr create_realtime_decoder( CUDA_QEC_INFO("Creating decoder {} of type {}", decoder_config.id, decoder_config.type); - auto pcm = cudaq::qec::pcm_from_sparse_vec(decoder_config.H_sparse, - decoder_config.syndrome_size, - decoder_config.block_size); - const auto num_observables = std::count(decoder_config.O_sparse.begin(), - decoder_config.O_sparse.end(), -1); - // Materialize O before decoder construction to validate its sparse shape and - // column indices for every decoder type. TRT also receives this matrix in its - // constructor parameters through prepare_decoder_params() below. - (void)cudaq::qec::pcm_from_sparse_vec( - decoder_config.O_sparse, num_observables, decoder_config.block_size); - auto decoder = cudaq::qec::get_decoder( - decoder_config.type, pcm, prepare_decoder_params(decoder_config)); + auto params = prepare_decoder_params(decoder_config); + std::unique_ptr decoder; + if (!decoder_config.stim_dem_path.empty()) { + // A DEM-native decoder is constructed from the model text itself. It + // derives its own dimensions and, like chromobius, installs the observable + // mapping its results are expressed in, so neither is supplied here. + decoder = cudaq::qec::get_decoder(decoder_config.type, + read_stim_dem(decoder_config), params); + const auto num_D_rows = std::count(decoder_config.D_sparse.begin(), + decoder_config.D_sparse.end(), -1); + if (num_D_rows != static_cast(decoder->get_syndrome_size())) + throw std::runtime_error(fmt::format( + "Number of rows in D_sparse vector is not equal to the number of " + "detectors in {}: {} != {}", + decoder_config.stim_dem_path, num_D_rows, + decoder->get_syndrome_size())); + } else { + auto pcm = cudaq::qec::pcm_from_sparse_vec(decoder_config.H_sparse, + decoder_config.syndrome_size, + decoder_config.block_size); + const auto num_observables = std::count(decoder_config.O_sparse.begin(), + decoder_config.O_sparse.end(), -1); + // Materialize O before decoder construction to validate its sparse shape + // and column indices for every decoder type. TRT also receives this matrix + // in its constructor parameters through prepare_decoder_params() above. + (void)cudaq::qec::pcm_from_sparse_vec( + decoder_config.O_sparse, num_observables, decoder_config.block_size); + decoder = cudaq::qec::get_decoder(decoder_config.type, pcm, params); + } decoder->set_decoder_id(decoder_config.id); - decoder->set_O_sparse(decoder_config.O_sparse); + if (!decoder_config.O_sparse.empty()) + decoder->set_O_sparse(decoder_config.O_sparse); decoder->set_D_sparse(decoder_config.D_sparse); // Force plugin initialization before the caller publishes the decoder for // realtime work. This preserves configure_decoders()'s existing behavior. auto t1 = std::chrono::high_resolution_clock::now(); - std::vector syndrome(decoder_config.syndrome_size, 0.0); + std::vector syndrome(decoder->get_syndrome_size(), 0.0); decoder->decode(syndrome); auto t2 = std::chrono::high_resolution_clock::now(); std::chrono::duration creation_duration = t1 - t0; diff --git a/libs/qec/python/bindings/py_decoding_config.cpp b/libs/qec/python/bindings/py_decoding_config.cpp index fb66c0832..27a0596ae 100644 --- a/libs/qec/python/bindings/py_decoding_config.cpp +++ b/libs/qec/python/bindings/py_decoding_config.cpp @@ -187,6 +187,7 @@ void bindDecodingConfig(nb::module_ &mod) { .def_rw("type", &decoder_config::type) .def_rw("dispatch", &decoder_config::dispatch) .def_rw("cuda_device_id", &decoder_config::cuda_device_id) + .def_rw("stim_dem_path", &decoder_config::stim_dem_path) .def_rw("block_size", &decoder_config::block_size) .def_rw("syndrome_size", &decoder_config::syndrome_size) .def_rw("H_sparse", &decoder_config::H_sparse) diff --git a/libs/qec/unittests/CMakeLists.txt b/libs/qec/unittests/CMakeLists.txt index 886ede5e7..471b5d623 100644 --- a/libs/qec/unittests/CMakeLists.txt +++ b/libs/qec/unittests/CMakeLists.txt @@ -53,6 +53,16 @@ if(TARGET cudaq-qec-decoding-server) cudaq-qec-realtime-decoding cudaq-qec-realtime-decoding-simulation cudaq::cudaq) + # The DEM-source test constructs chromobius through the plugin loader. + if(TARGET cudaq-qec-chromobius) + add_dependencies(test_decoders_yaml cudaq-qec-chromobius) + target_compile_definitions(test_decoders_yaml PRIVATE CUDAQX_QEC_HAS_CHROMOBIUS) + endif() + # The DEM-source test constructs chromobius through the plugin loader. + if(TARGET cudaq-qec-chromobius) + add_dependencies(test_decoders_yaml cudaq-qec-chromobius) + target_compile_definitions(test_decoders_yaml PRIVATE CUDAQX_QEC_HAS_CHROMOBIUS) + endif() add_dependencies(CUDAQXQECUnitTests test_decoders_yaml) gtest_discover_tests(test_decoders_yaml) diff --git a/libs/qec/unittests/test_decoders_yaml.cpp b/libs/qec/unittests/test_decoders_yaml.cpp index 111f9e588..08e6c2f91 100644 --- a/libs/qec/unittests/test_decoders_yaml.cpp +++ b/libs/qec/unittests/test_decoders_yaml.cpp @@ -1310,3 +1310,98 @@ TEST(DecoderYAMLTest, NonSchemaKeysDroppedFromDecoderParamsAndEmission) { // The stored args are untouched -- only the derived views are filtered. EXPECT_TRUE(config.decoder_custom_args.map().contains("not_a_real_param")); } + +#ifdef CUDAQX_QEC_HAS_CHROMOBIUS +// Four detectors carrying the colour/basis annotation chromobius decodes from, +// matching the model used by the chromobius unit tests. +constexpr const char *kChromobiusDem = R"DEM( +error(0.1) D0 D1 +error(0.1) D0 D1 D2 +error(0.1) D0 L0 +error(0.1) D1 D2 D3 +error(0.1) D2 D3 +error(0.1) D3 +detector(0, 0, 0, 1) D0 +detector(1, 0, 0, 2) D1 +detector(2, 0, 0, 0) D2 +detector(3, 0, 0, 1) D3 +)DEM"; + +/// Writes the model to a file, since a config names a DEM by path. +class ScopedDemFile { +public: + ScopedDemFile() { + static int counter = 0; + path_ = std::filesystem::temp_directory_path() / + ("unblock-" + std::to_string(getpid()) + "-" + + std::to_string(counter++) + ".dem"); + std::ofstream(path_) << kChromobiusDem; + } + ~ScopedDemFile() { + std::error_code ec; + std::filesystem::remove(path_, ec); + } + std::string path() const { return path_.string(); } + +private: + std::filesystem::path path_; +}; + +TEST(ChromobiusOnDecodingServer, ConstructsFromARawDemSource) { + ScopedDemFile dem; + + // No H_sparse and no O_sparse: the DEM is the model, and chromobius installs + // the observable mapping its results are expressed in. + cudaq::qec::decoding::config::decoder_config dc; + dc.id = 0; + dc.type = "chromobius"; + dc.stim_dem_path = dem.path(); + dc.D_sparse = {0, -1, 1, -1, 2, -1, 3, -1}; + + cudaq::qec::decoding::config::multi_decoder_config mc; + mc.decoders.push_back(dc); + ASSERT_EQ(cudaq::qec::decoding::config::configure_decoders(mc), 0); + cudaq::qec::decoding::config::finalize_decoders(); +} + +TEST(ChromobiusOnDecodingServer, MatrixOnlyConfigStillFails) { + // The converse: without a DEM source chromobius cannot be built, which is + // what naming a DEM lifts. If this ever passes, the test above proves + // nothing. + auto dc = create_test_sample_realtime_decoder_config(0); + dc.type = "chromobius"; + + cudaq::qec::decoding::config::multi_decoder_config mc; + mc.decoders.push_back(dc); + EXPECT_NE(cudaq::qec::decoding::config::configure_decoders(mc), 0); + cudaq::qec::decoding::config::finalize_decoders(); +} + +TEST(ChromobiusOnDecodingServer, DemReachesANestedGlobalDecoder) { + ScopedDemFile dem; + + // The nesting half: a DEM-sourced entry hands the model text down through + // global_decoder_params, which is where trt_decoder builds chromobius from. + cudaq::qec::decoding::config::decoder_config dc; + dc.id = 0; + dc.type = "trt_decoder"; + dc.stim_dem_path = dem.path(); + dc.D_sparse = {0, -1, 1, -1, 2, -1, 3, -1}; + dc.decoder_custom_args.map().insert("global_decoder", + std::string{"chromobius"}); + dc.decoder_custom_args.map().insert("global_decoder_params", + cudaqx::heterogeneous_map{}); + + auto params = cudaq::qec::decoding::host::prepare_decoder_params(dc); + ASSERT_TRUE(params.contains("global_decoder_params")); + const auto global_params = + params.get("global_decoder_params"); + ASSERT_TRUE(global_params.contains("stim_dem")); + + // What trt_decoder passes to decoder::get() must build the decoder the + // config named. + auto global_decoder = cudaq::qec::decoder::get( + "chromobius", global_params.get("stim_dem"), global_params); + EXPECT_EQ(global_decoder->get_syndrome_size(), 4); +} +#endif // CUDAQX_QEC_HAS_CHROMOBIUS