From 07c5670f914b09b5abe47a3b2ef1bcfc76055f9b Mon Sep 17 00:00:00 2001 From: Carsten Burgard Date: Fri, 15 Dec 2023 16:25:47 +0100 Subject: [PATCH 1/6] [RF] Error message in ParamHistFunc::getParameter() for invalid index --- roofit/histfactory/src/ParamHistFunc.cxx | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/roofit/histfactory/src/ParamHistFunc.cxx b/roofit/histfactory/src/ParamHistFunc.cxx index 29341ab68402f..cf10bb915fe53 100644 --- a/roofit/histfactory/src/ParamHistFunc.cxx +++ b/roofit/histfactory/src/ParamHistFunc.cxx @@ -208,7 +208,11 @@ RooAbsReal& ParamHistFunc::getParameter( Int_t index ) const { const int j = tmp / n.z; const int k = tmp % n.z; - return static_cast(_paramSet[i + j * n.x + k * n.xy]); + const int idx = i + j * n.x + k * n.xy; + if (idx >= _numBins) { + throw std::runtime_error("invalid index"); + } + return static_cast(_paramSet[idx]); } From cdcd574cfbcbe06a4f611552df644be3dfd4a6ca Mon Sep 17 00:00:00 2001 From: Carsten Burgard Date: Wed, 20 Dec 2023 21:23:31 +0100 Subject: [PATCH 2/6] [RF][HS3] Added importer for `normal_dist` Gaussians can now be imported with "normal_dist" as well, as described by the standard. --- etc/RooFitHS3_wsfactoryexpressions.json | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/etc/RooFitHS3_wsfactoryexpressions.json b/etc/RooFitHS3_wsfactoryexpressions.json index e09474b0493ee..08953c21c4d02 100644 --- a/etc/RooFitHS3_wsfactoryexpressions.json +++ b/etc/RooFitHS3_wsfactoryexpressions.json @@ -51,6 +51,14 @@ "sigma" ] }, + "normal_dist": { + "class": "RooGaussian", + "arguments": [ + "x", + "mean", + "sigma" + ] + }, "interpolation0d": { "class": "RooStats::HistFactory::FlexibleInterpVar", "arguments": [ From 49ccf701201509ec8d772473de5a1020e8dfccac Mon Sep 17 00:00:00 2001 From: Carsten Burgard Date: Wed, 20 Dec 2023 21:24:12 +0100 Subject: [PATCH 3/6] [RF][HS3] Improved error reporting from JSON parsing errors Errors from the nlohmann_json importer are now correctly forwarded to the user, rather than giving an unspecific "std_exception". --- roofit/jsoninterface/src/JSONParser.cxx | 23 ++++++++++++++++------- roofit/jsoninterface/src/JSONParser.h | 2 +- 2 files changed, 17 insertions(+), 8 deletions(-) diff --git a/roofit/jsoninterface/src/JSONParser.cxx b/roofit/jsoninterface/src/JSONParser.cxx index fc6e927d0f99f..9341f51090c36 100644 --- a/roofit/jsoninterface/src/JSONParser.cxx +++ b/roofit/jsoninterface/src/JSONParser.cxx @@ -16,6 +16,17 @@ #include +namespace { +inline nlohmann::json parseWrapper(std::istream &is) +{ + try { + return nlohmann::json::parse(is); + } catch (const nlohmann::json::exception &ex) { + throw std::runtime_error(ex.what()); + } +} +} // namespace + // TJSONTree methods TJSONTree::TJSONTree() : root(this){}; @@ -70,7 +81,7 @@ class TJSONTree::Node::Impl::BaseNode : public TJSONTree::Node::Impl { public: nlohmann::json &get() override { return node; } const nlohmann::json &get() const override { return node; } - BaseNode(std::istream &is) : Impl(""), node(nlohmann::json::parse(is)) {} + BaseNode(std::istream &is) : Impl(""), node(parseWrapper(is)) {} BaseNode() : Impl("") {} }; @@ -108,8 +119,6 @@ TJSONTree::Node::Node(TJSONTree *t, Impl &other) TJSONTree::Node::Node(const Node &other) : Node(other.tree, *other.node) {} -TJSONTree::Node::~Node() {} - // TJSONNode interface void TJSONTree::Node::writeJSON(std::ostream &os) const @@ -200,7 +209,7 @@ TJSONTree::Node &TJSONTree::Node::set_map() if (isResettingPossible(node->get())) { node->get() = nlohmann::json::object(); } else { - throw std::runtime_error("cannot declare " + this->key() + " to be of map-type, already of type " + + throw std::runtime_error("cannot declare \"" + this->key() + "\" to be of map - type, already of type " + node->get().type_name()); } return *this; @@ -214,7 +223,7 @@ TJSONTree::Node &TJSONTree::Node::set_seq() if (isResettingPossible(node->get())) { node->get() = nlohmann::json::array(); } else { - throw std::runtime_error("cannot declare " + this->key() + " to be of seq-type, already of type " + + throw std::runtime_error("cannot declare \"" + this->key() + "\" to be of seq - type, already of type " + node->get().type_name()); } return *this; @@ -239,8 +248,8 @@ std::string TJSONTree::Node::val() const case nlohmann::json::value_t::number_unsigned: return std::to_string(node->get().get()); case nlohmann::json::value_t::number_float: return std::to_string(node->get().get()); default: - throw std::runtime_error(std::string("node " + node->key() + ": implicit string conversion for type " + - node->get().type_name() + " not supported!")); + throw std::runtime_error("node \"" + node->key() + "\": implicit string conversion for type " + + node->get().type_name() + " not supported!"); } } diff --git a/roofit/jsoninterface/src/JSONParser.h b/roofit/jsoninterface/src/JSONParser.h index d05625df6657e..06270868598d4 100644 --- a/roofit/jsoninterface/src/JSONParser.h +++ b/roofit/jsoninterface/src/JSONParser.h @@ -42,7 +42,7 @@ class TJSONTree : public RooFit::Detail::JSONTree { Node(TJSONTree *t, Impl &other); Node(TJSONTree *t); Node(const Node &other); - virtual ~Node(); + virtual ~Node() = default; Node &operator<<(std::string const &s) override; Node &operator<<(int i) override; Node &operator<<(double d) override; From b2b57cafda9bd9d97cfc43a71ab109f449518863 Mon Sep 17 00:00:00 2001 From: Carsten Burgard Date: Wed, 20 Dec 2023 21:25:06 +0100 Subject: [PATCH 4/6] [RF][HS3] Improved robustness, removed warnings The implementation is a bit more lenient when missing some values (e.g. ignoring cases where the histfactory modifier name has been omitted and instead just numbers the systematics) --- roofit/hs3/src/JSONFactories_HistFactory.cxx | 13 ++++++++----- 1 file changed, 8 insertions(+), 5 deletions(-) diff --git a/roofit/hs3/src/JSONFactories_HistFactory.cxx b/roofit/hs3/src/JSONFactories_HistFactory.cxx index 09d56e33a7650..e090d620400d7 100644 --- a/roofit/hs3/src/JSONFactories_HistFactory.cxx +++ b/roofit/hs3/src/JSONFactories_HistFactory.cxx @@ -274,9 +274,14 @@ bool importHistSample(RooJSONFactoryWSTool &tool, RooDataHist &dh, RooArgSet con RooArgList histoLo; RooArgList histoHi; + int idx = 0; for (const auto &mod : p["modifiers"].children()) { std::string const &modtype = mod["type"].val(); - std::string const &sysname = mod["name"].val(); + std::string const &sysname = + mod.has_child("name") + ? mod["name"].val() + : (mod.has_child("parameter") ? mod["parameter"].val() : "syst_" + std::to_string(idx)); + ++idx; if (modtype == "staterror") { // this is dealt with at a different place, ignore it for now } else if (modtype == "normfactor") { @@ -925,12 +930,10 @@ bool tryExportHistFactory(RooJSONFactoryWSTool *tool, const std::string &pdfname optionallyExportGammaParameters(mod, sys.name, sys.parameters); mod["constraint"] << toString(sys.constraint); if (sys.constraint) { - auto &data = mod["data"].set_map(); - auto &vals = data["vals"]; + auto &vals = mod["data"].set_map()["vals"]; vals.fill_seq(sys.constraints); } else { - auto &data = mod["data"].set_map(); - auto &vals = data["vals"]; + auto &vals = mod["data"].set_map()["vals"]; vals.set_seq(); for (std::size_t i = 0; i < sys.parameters.size(); ++i) { vals.append_child() << 0; From a726681a534c77f488c9001f92c6e81e5d2489ad Mon Sep 17 00:00:00 2001 From: Carsten Burgard Date: Wed, 20 Dec 2023 21:25:36 +0100 Subject: [PATCH 5/6] [RF][HS3] Add method to fill workspace with all variables in "domains" --- roofit/hs3/src/Domains.cxx | 36 ++++++++++++++++++++++++++++++++---- roofit/hs3/src/Domains.h | 5 +++++ 2 files changed, 37 insertions(+), 4 deletions(-) diff --git a/roofit/hs3/src/Domains.cxx b/roofit/hs3/src/Domains.cxx index 3eb3a7d51cb92..57643764daf1a 100644 --- a/roofit/hs3/src/Domains.cxx +++ b/roofit/hs3/src/Domains.cxx @@ -15,6 +15,7 @@ #include #include #include +#include #include @@ -22,9 +23,18 @@ namespace RooFit { namespace JSONIO { namespace Detail { +constexpr static auto defaultDomainName = "default_domain"; + +void Domains::populate(RooWorkspace &ws) const +{ + auto found = _map.find(defaultDomainName); + if (found != _map.end()) { + found->second.populate(ws); + } +} void Domains::readVariable(const char *name, double min, double max) { - _map["default_domain"].readVariable(name, min, max); + _map[defaultDomainName].readVariable(name, min, max); } void Domains::readVariable(RooRealVar const &var) { @@ -32,12 +42,16 @@ void Domains::readVariable(RooRealVar const &var) } void Domains::writeVariable(RooRealVar &var) const { - _map.at("default_domain").writeVariable(var); + _map.at(defaultDomainName).writeVariable(var); } void Domains::readJSON(RooFit::Detail::JSONNode const &node) { - _map["default_domain"].readJSON(*RooJSONFactoryWSTool::findNamedChild(node, "default_domain")); + auto defaultDomain = RooJSONFactoryWSTool::findNamedChild(node, defaultDomainName); + if (!defaultDomain) { + RooJSONFactoryWSTool::error("\"domains\" do not contain \"" + std::string{defaultDomainName} + "\""); + } + _map[defaultDomainName].readJSON(*defaultDomain); } void Domains::writeJSON(RooFit::Detail::JSONNode &node) const { @@ -71,7 +85,9 @@ void Domains::ProductDomain::writeVariable(RooRealVar &var) const } void Domains::ProductDomain::readJSON(RooFit::Detail::JSONNode const &node) { - // In the future, throw an exception if the type is not product domain + if (!node.has_child("type") || node["type"].val() != "product_domain") { + RooJSONFactoryWSTool::error("only domains of type \"product_domain\" are currently supported!"); + } for (auto const &varNode : node["axes"].children()) { auto &elem = _map[RooJSONFactoryWSTool::name(varNode)]; @@ -101,6 +117,18 @@ void Domains::ProductDomain::writeJSON(RooFit::Detail::JSONNode &node) const varnode["max"] << elem.max; } } +void Domains::ProductDomain::populate(RooWorkspace &ws) const +{ + for (auto const &item : _map) { + const auto &name = item.first; + if (!ws.var(name)) { + const auto &elem = item.second; + const double vMin = elem.hasMin ? elem.min : -RooNumber::infinity(); + const double vMax = elem.hasMax ? elem.max : RooNumber::infinity(); + ws.import(RooRealVar{name.c_str(), name.c_str(), vMin, vMax}); + } + } +} } // namespace Detail } // namespace JSONIO diff --git a/roofit/hs3/src/Domains.h b/roofit/hs3/src/Domains.h index 6d542dbe82cef..0f4f14acff189 100644 --- a/roofit/hs3/src/Domains.h +++ b/roofit/hs3/src/Domains.h @@ -18,6 +18,7 @@ #include class RooRealVar; +class RooWorkspace; namespace RooFit { namespace Detail { @@ -39,6 +40,8 @@ class Domains { void readJSON(RooFit::Detail::JSONNode const &); void writeJSON(RooFit::Detail::JSONNode &) const; + void populate(RooWorkspace &ws) const; + private: class ProductDomain { public: @@ -48,6 +51,8 @@ class Domains { void readJSON(RooFit::Detail::JSONNode const &); void writeJSON(RooFit::Detail::JSONNode &) const; + void populate(RooWorkspace &ws) const; + private: struct ProductDomainElement { bool hasMin = false; From b7d87f83b0db94ac269c8dd8b51e58ad144afc8b Mon Sep 17 00:00:00 2001 From: Carsten Burgard Date: Wed, 20 Dec 2023 21:26:01 +0100 Subject: [PATCH 6/6] [RF][HS3] Improved error reporting, using new populate feature The list of variables in the workspace is now no longer just filled from the parameter_points, but also from the domains, allowing cases where no parameter points are given to be imported successfully. Also, the HS3 version tag has been correctly updated to 0.2. --- roofit/hs3/src/RooJSONFactoryWSTool.cxx | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/roofit/hs3/src/RooJSONFactoryWSTool.cxx b/roofit/hs3/src/RooJSONFactoryWSTool.cxx index c84e5e22c0e85..3faa615664c7b 100644 --- a/roofit/hs3/src/RooJSONFactoryWSTool.cxx +++ b/roofit/hs3/src/RooJSONFactoryWSTool.cxx @@ -93,7 +93,7 @@ tool.writedoc("hs3.tex") ~~~ */ -constexpr auto hs3VersionTag = "0.1.90"; +constexpr auto hs3VersionTag = "0.2"; using RooFit::Detail::JSONNode; using RooFit::Detail::JSONTree; @@ -207,6 +207,9 @@ bool isValidName(const std::string &str) */ void configureVariable(RooFit::JSONIO::Detail::Domains &domains, const JSONNode &p, RooRealVar &v) { + if (!p.has_child("name")) { + RooJSONFactoryWSTool::error("cannot instantiate variable without \"name\"!"); + } if (auto n = p.find("value")) v.setVal(n->val_double()); domains.writeVariable(v); @@ -2072,6 +2075,7 @@ void RooJSONFactoryWSTool::importAllNodes(const JSONNode &n) if (auto domains = n.find("domains")) { _domains->readJSON(*domains); } + _domains->populate(_workspace); _rootnodeInput = &n;