From 68c6e2b0d52737fd753a3e8915778ffc51137222 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?David--Cl=C3=A9ris=20Timoth=C3=A9e?= Date: Thu, 20 Aug 2026 10:03:32 +0000 Subject: [PATCH 1/2] [Patch] Out-of-line FieldVariant visit, PatchDataLayer, and add_field Forward FieldVariant visitors to std::visit without wrapping lambdas, move PatchDataLayer copy/visit helpers into the .cpp, and extern-template PatchDataLayerLayout::add_field for the 18 enabled field types. Assisted-by: Cursor --- .../include/shamrock/patch/FieldVariant.hpp | 38 ++---- .../include/shamrock/patch/PatchDataLayer.hpp | 114 ++---------------- .../shamrock/patch/PatchDataLayerLayout.hpp | 32 ++--- src/shamrock/src/patch/PatchDataLayer.cpp | 102 ++++++++++++++++ .../src/patch/PatchDataLayerLayout.cpp | 29 +++++ .../patch/PatchDataLayerLayoutTests.cpp | 18 +++ src/tests/shamrock/patch/PatchDataTests.cpp | 28 +++++ 7 files changed, 211 insertions(+), 150 deletions(-) diff --git a/src/shamrock/include/shamrock/patch/FieldVariant.hpp b/src/shamrock/include/shamrock/patch/FieldVariant.hpp index a994bcfcff..b697c174bc 100644 --- a/src/shamrock/include/shamrock/patch/FieldVariant.hpp +++ b/src/shamrock/include/shamrock/patch/FieldVariant.hpp @@ -18,6 +18,7 @@ #include "shambase/exception.hpp" #include "shambackends/typeAliasVec.hpp" +#include #include namespace shamrock::patch { @@ -65,49 +66,36 @@ namespace shamrock::patch { "the type asked is not correct"); } + /** + * @brief Apply a visitor to the held alternative. + * + * The functor is forwarded to std::visit unchanged. Wrapping it in an extra + * generic lambda would give every call site a unique visitor type, and + * libstdc++ builds a visit vtable for each of those types. + */ template void visit(Func &&f) { - std::visit( - [&](auto &arg) { - f(arg); - }, - value); + std::visit(std::forward(f), value); } template auto visit_return(Func &&f) { - return std::visit( - [&](auto &arg) { - return f(arg); - }, - value); + return std::visit(std::forward(f), value); } template void visit(Func &&f) const { - std::visit( - [&](auto &arg) { - f(arg); - }, - value); + std::visit(std::forward(f), value); } template auto visit_return(Func &&f) const { - return std::visit( - [&](auto &arg) { - return f(arg); - }, - value); + return std::visit(std::forward(f), value); } template class Container2, class Func> FieldVariant convert(Func &&f) { - return std::visit( - [&](auto &arg) { - return f(arg); - }, - value); + return std::visit(std::forward(f), value); } }; diff --git a/src/shamrock/include/shamrock/patch/PatchDataLayer.hpp b/src/shamrock/include/shamrock/patch/PatchDataLayer.hpp index 641b65bef5..792639dcce 100644 --- a/src/shamrock/include/shamrock/patch/PatchDataLayer.hpp +++ b/src/shamrock/include/shamrock/patch/PatchDataLayer.hpp @@ -79,19 +79,7 @@ namespace shamrock::patch { init_fields(); } - inline PatchDataLayer(const PatchDataLayer &other) : pdl_ptr(other.get_layout_ptr()) { - - NamedStackEntry stack_loc{"PatchDataLayer::copy_constructor", true}; - - for (auto &field_var : other.fields) { - - field_var.visit([&](auto &field) { - using base_t = - typename std::remove_reference::type::Field_type; - fields.emplace_back(PatchDataField(field)); - }); - }; - } + PatchDataLayer(const PatchDataLayer &other); /** * @brief PatchDataLayer move constructor @@ -120,18 +108,14 @@ namespace shamrock::patch { template inline void for_each_field_any(Functor &&func) { for (auto &f : fields) { - f.visit([&](auto &arg) { - func(arg); - }); + f.visit(func); } } template inline void for_each_field_any(Functor &&func) const { for (auto &f : fields) { - f.visit([&](const auto &arg) { - func(arg); - }); + f.visit(func); } } @@ -234,42 +218,13 @@ namespace shamrock::patch { void append_subset_to( const sham::DeviceBuffer &idxs_buf, u32 sz, PatchDataLayer &pdat) const; - inline u32 get_obj_cnt() const { + u32 get_obj_cnt() const; - bool is_empty = fields.empty(); - - if (!is_empty) { - return fields[0].visit_return([](const auto &field) { - return field.get_obj_cnt(); - }); - } - - throw shambase::make_except_with_loc( - "this PatchDataLayer does not contain any fields"); - } - - inline u64 memsize() { - u64 sum = 0; - - for (auto &field_var : fields) { - - field_var.visit([&](auto &field) { - sum += field.memsize(); - }); - } - - return sum; - } + u64 memsize(); inline bool is_empty() { return get_obj_cnt() == 0; } - void synchronize_buf() { - for (auto &field_var : fields) { - field_var.visit([&](auto &field) { - field.synchronize_buf(); - }); - } - } + void synchronize_buf(); void overwrite(PatchDataLayer &pdat, u32 obj_cnt); @@ -389,17 +344,7 @@ namespace shamrock::patch { * @brief check that all contained field have the same obj cnt * */ - inline void check_field_obj_cnt_match() { - u32 cnt = get_obj_cnt(); - for (auto &field_var : fields) { - field_var.visit([&](auto &field) { - if (field.get_obj_cnt() != cnt) { - throw shambase::make_except_with_loc( - "mismatch in obj cnt"); - } - }); - } - } + void check_field_obj_cnt_match(); // template inline std::vector & > get_field_list(){ // std::vector & > ret; @@ -431,48 +376,9 @@ namespace shamrock::patch { void fields_raz(); - bool has_nan() { - StackEntry stack_loc{}; - - bool ret = false; - - for (auto &field_var : fields) { - field_var.visit([&](auto &field) { - if (field.has_nan()) { - ret = true; - } - }); - } - return ret; - } - bool has_inf() { - StackEntry stack_loc{}; - - bool ret = false; - - for (auto &field_var : fields) { - field_var.visit([&](auto &field) { - if (field.has_inf()) { - ret = true; - } - }); - } - return ret; - } - bool has_nan_or_inf() { - StackEntry stack_loc{}; - - bool ret = false; - - for (auto &field_var : fields) { - field_var.visit([&](auto &field) { - if (field.has_nan_or_inf()) { - ret = true; - } - }); - } - return ret; - } + bool has_nan(); + bool has_inf(); + bool has_nan_or_inf(); /** * @brief diff --git a/src/shamrock/include/shamrock/patch/PatchDataLayerLayout.hpp b/src/shamrock/include/shamrock/patch/PatchDataLayerLayout.hpp index 3ecc072c08..96fed6b3a8 100644 --- a/src/shamrock/include/shamrock/patch/PatchDataLayerLayout.hpp +++ b/src/shamrock/include/shamrock/patch/PatchDataLayerLayout.hpp @@ -20,6 +20,7 @@ #include "shambase/exception.hpp" #include "shambase/string.hpp" #include "nlohmann/json_fwd.hpp" +#include "shamrock/legacy/patch/base/enabled_fields.hpp" #include "shamrock/patch/FieldVariant.hpp" #include "shamsys/legacy/log.hpp" #include @@ -187,9 +188,7 @@ namespace shamrock::patch { template inline void for_each_field_any(Functor &&func) const { for (auto &f : fields) { - f.visit([&](auto &arg) { - func(arg); - }); + f.visit(func); } } @@ -296,24 +295,15 @@ namespace shamrock::patch { // out of line implementation of the PatchDataLayerLayout //////////////////////////////////////////////////////////////////////////////////////////////// - template - inline void PatchDataLayerLayout::add_field( - const std::string &field_name, u32 nvar, SourceLocation loc) { - if (has_field_name(field_name)) { - throw shambase::make_except_with_loc( - "add_field -> the name already exists"); - } - - shamlog_debug_ln( - "PatchDataLayerLayout", - "adding field :", - field_name, - nvar, - "loc :", - loc.format_one_line()); - - fields.push_back(var_t{FieldDescriptor(field_name, nvar)}); - } +#ifndef DOXYGEN + // Explicit instantiations live in PatchDataLayerLayout.cpp so TUs that call add_field + // do not instantiate FieldVariant construction for every enabled field type. + #define X(type) \ + extern template void PatchDataLayerLayout::add_field( \ + const std::string &field_name, u32 nvar, SourceLocation loc); + XMAC_LIST_ENABLED_FIELD + #undef X +#endif template inline PatchDataLayerLayout::FieldDescriptor PatchDataLayerLayout::get_field( diff --git a/src/shamrock/src/patch/PatchDataLayer.cpp b/src/shamrock/src/patch/PatchDataLayer.cpp index f6f162a1fa..9d20b08caf 100644 --- a/src/shamrock/src/patch/PatchDataLayer.cpp +++ b/src/shamrock/src/patch/PatchDataLayer.cpp @@ -57,6 +57,108 @@ namespace shamrock::patch { }); } + PatchDataLayer::PatchDataLayer(const PatchDataLayer &other) : pdl_ptr(other.get_layout_ptr()) { + + NamedStackEntry stack_loc{"PatchDataLayer::copy_constructor", true}; + + for (auto &field_var : other.fields) { + + field_var.visit([&](auto &field) { + using base_t = typename std::remove_reference::type::Field_type; + fields.emplace_back(PatchDataField(field)); + }); + } + } + + u32 PatchDataLayer::get_obj_cnt() const { + + if (!fields.empty()) { + return fields[0].visit_return([](const auto &field) { + return field.get_obj_cnt(); + }); + } + + throw shambase::make_except_with_loc( + "this PatchDataLayer does not contain any fields"); + } + + u64 PatchDataLayer::memsize() { + u64 sum = 0; + + for (auto &field_var : fields) { + + field_var.visit([&](auto &field) { + sum += field.memsize(); + }); + } + + return sum; + } + + void PatchDataLayer::synchronize_buf() { + for (auto &field_var : fields) { + field_var.visit([&](auto &field) { + field.synchronize_buf(); + }); + } + } + + void PatchDataLayer::check_field_obj_cnt_match() { + u32 cnt = get_obj_cnt(); + for (auto &field_var : fields) { + field_var.visit([&](auto &field) { + if (field.get_obj_cnt() != cnt) { + throw shambase::make_except_with_loc("mismatch in obj cnt"); + } + }); + } + } + + bool PatchDataLayer::has_nan() { + StackEntry stack_loc{}; + + bool ret = false; + + for (auto &field_var : fields) { + field_var.visit([&](auto &field) { + if (field.has_nan()) { + ret = true; + } + }); + } + return ret; + } + + bool PatchDataLayer::has_inf() { + StackEntry stack_loc{}; + + bool ret = false; + + for (auto &field_var : fields) { + field_var.visit([&](auto &field) { + if (field.has_inf()) { + ret = true; + } + }); + } + return ret; + } + + bool PatchDataLayer::has_nan_or_inf() { + StackEntry stack_loc{}; + + bool ret = false; + + for (auto &field_var : fields) { + field_var.visit([&](auto &field) { + if (field.has_nan_or_inf()) { + ret = true; + } + }); + } + return ret; + } + void PatchDataLayer::extract_element(u32 pidx, PatchDataLayer &out_pdat) { StackEntry stack_loc{}; diff --git a/src/shamrock/src/patch/PatchDataLayerLayout.cpp b/src/shamrock/src/patch/PatchDataLayerLayout.cpp index 55d118c93a..3bf155056c 100644 --- a/src/shamrock/src/patch/PatchDataLayerLayout.cpp +++ b/src/shamrock/src/patch/PatchDataLayerLayout.cpp @@ -15,6 +15,8 @@ */ #include "shamrock/patch/PatchDataLayerLayout.hpp" +#include "shamrock/legacy/patch/base/enabled_fields.hpp" +#include "shamsys/legacy/log.hpp" #include namespace shamrock::patch { @@ -197,4 +199,31 @@ namespace shamrock::patch { return ret; } + template + void PatchDataLayerLayout::add_field( + const std::string &field_name, u32 nvar, SourceLocation loc) { + if (has_field_name(field_name)) { + throw shambase::make_except_with_loc( + "add_field -> the name already exists"); + } + + shamlog_debug_ln( + "PatchDataLayerLayout", + "adding field :", + field_name, + nvar, + "loc :", + loc.format_one_line()); + + fields.push_back(var_t{FieldDescriptor(field_name, nvar)}); + } + +#ifndef DOXYGEN + #define X(type) \ + template void PatchDataLayerLayout::add_field( \ + const std::string &field_name, u32 nvar, SourceLocation loc); + XMAC_LIST_ENABLED_FIELD + #undef X +#endif + } // namespace shamrock::patch diff --git a/src/tests/shamrock/patch/PatchDataLayerLayoutTests.cpp b/src/tests/shamrock/patch/PatchDataLayerLayoutTests.cpp index 2aa85ad2fe..98717edaad 100644 --- a/src/tests/shamrock/patch/PatchDataLayerLayoutTests.cpp +++ b/src/tests/shamrock/patch/PatchDataLayerLayoutTests.cpp @@ -7,6 +7,7 @@ // // -------------------------------------------------------// +#include "shamrock/legacy/patch/base/enabled_fields.hpp" #include "shamrock/patch/PatchDataLayerLayout.hpp" #include "shamtest/shamtest.hpp" #include @@ -44,3 +45,20 @@ NEW_TEST(Unittest, "shamrock/patch/PatchDataLayerLayout::serialize_json", 1) { REQUIRE(pdl == pdl_out); } + +NEW_TEST(Unittest, "shamrock/patch/PatchDataLayerLayout::add_field", 1) { + using namespace shamrock::patch; + + PatchDataLayerLayout pdl; + + u32 nfields = 0; +#define X(type) \ + pdl.add_field(#type, 1); \ + nfields++; + XMAC_LIST_ENABLED_FIELD +#undef X + + REQUIRE_EQUAL(pdl.get_field_names().size(), nfields); + + REQUIRE_EXCEPTION_THROW(pdl.add_field("f32", 1), std::invalid_argument); +} diff --git a/src/tests/shamrock/patch/PatchDataTests.cpp b/src/tests/shamrock/patch/PatchDataTests.cpp index 2cc939a603..f378215681 100644 --- a/src/tests/shamrock/patch/PatchDataTests.cpp +++ b/src/tests/shamrock/patch/PatchDataTests.cpp @@ -285,3 +285,31 @@ NEW_TEST(Unittest, "shamrock/patch/PatchDataLayer::operator==", 1) { REQUIRE_NAMED("object count mismatch", !(a == b)); } } + +NEW_TEST(Unittest, "shamrock/patch/PatchDataLayer::copy_constructor", 1) { + using namespace shamrock::patch; + + constexpr u32 obj_cnt = 16; + constexpr u64 seed = 0x222; + + std::shared_ptr pdl_ptr = std::make_shared(); + pdl_ptr->add_field("a", 1); + pdl_ptr->add_field("b", 2); + pdl_ptr->add_field("c", 1); + + PatchDataLayer a = PatchDataLayer::mock_patchdata(seed, obj_cnt, pdl_ptr); + PatchDataLayer b{a}; + PatchDataLayer c = a.duplicate(); + + REQUIRE_NAMED("copy equal", a == b); + REQUIRE_NAMED("duplicate equal", a == c); + REQUIRE_EQUAL(a.get_obj_cnt(), b.get_obj_cnt()); + REQUIRE_EQUAL(a.memsize(), b.memsize()); + REQUIRE_EQUAL(a.has_nan(), false); + REQUIRE_EQUAL(b.has_nan(), false); + REQUIRE_EQUAL(a.has_inf(), false); + REQUIRE_EQUAL(a.has_nan_or_inf(), false); + + a.check_field_obj_cnt_match(); + b.check_field_obj_cnt_match(); +} From b534ba58df414a76b3dc655e3ca801fc2be4a8cc Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?David--Cl=C3=A9ris=20Timoth=C3=A9e?= Date: Thu, 20 Aug 2026 11:15:24 +0000 Subject: [PATCH 2/2] [Patch] Fix copy_constructor test for mock inf values mock_patchdata samples the full type range, so fields can contain inf. Assert that copies preserve NaN/Inf instead of assuming finite mock data. Assisted-by: Cursor --- src/tests/shamrock/patch/PatchDataTests.cpp | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/src/tests/shamrock/patch/PatchDataTests.cpp b/src/tests/shamrock/patch/PatchDataTests.cpp index f378215681..98ad248a46 100644 --- a/src/tests/shamrock/patch/PatchDataTests.cpp +++ b/src/tests/shamrock/patch/PatchDataTests.cpp @@ -305,10 +305,12 @@ NEW_TEST(Unittest, "shamrock/patch/PatchDataLayer::copy_constructor", 1) { REQUIRE_NAMED("duplicate equal", a == c); REQUIRE_EQUAL(a.get_obj_cnt(), b.get_obj_cnt()); REQUIRE_EQUAL(a.memsize(), b.memsize()); - REQUIRE_EQUAL(a.has_nan(), false); - REQUIRE_EQUAL(b.has_nan(), false); - REQUIRE_EQUAL(a.has_inf(), false); - REQUIRE_EQUAL(a.has_nan_or_inf(), false); + // mock_patchdata samples the full type range, so fields can contain inf + REQUIRE_EQUAL(a.has_nan(), b.has_nan()); + REQUIRE_EQUAL(a.has_inf(), b.has_inf()); + REQUIRE_EQUAL(a.has_nan_or_inf(), b.has_nan_or_inf()); + REQUIRE_EQUAL(c.has_nan(), a.has_nan()); + REQUIRE_EQUAL(c.has_inf(), a.has_inf()); a.check_field_obj_cnt_match(); b.check_field_obj_cnt_match();