From 36ca05529754f0225700e038adf085ec6f758f50 Mon Sep 17 00:00:00 2001 From: Martin Date: Wed, 24 Jun 2026 16:21:40 -0500 Subject: [PATCH 1/9] Add retry mechanism and dynamic model loading to SONIC framework. Resolve cmsTriton conflict in CMSSW_17_0_0_pre2. Co-authored-by: Trevin Lee --- HeterogeneousCore/SonicCore/BuildFile.xml | 1 + .../SonicCore/interface/RetryActionBase.h | 36 + .../SonicCore/interface/SonicClientBase.h | 14 +- .../SonicCore/plugins/BuildFile.xml | 6 + .../plugins/RetrySameServerAction.cc | 28 + .../SonicCore/src/RetryActionBase.cc | 15 + .../SonicCore/src/SonicClientBase.cc | 77 +- .../SonicCore/test/DummyClient.h | 2 +- .../SonicCore/test/sonicTestAna_cfg.py | 15 +- .../SonicCore/test/sonicTest_cfg.py | 45 +- HeterogeneousCore/SonicTriton/BuildFile.xml | 3 +- .../interface/RetryActionDiffServer.h | 32 + .../interface/RetryFallbackServerAction.h | 18 + .../SonicTriton/interface/TritonClient.h | 222 +-- .../SonicTriton/interface/TritonService.h | 58 +- .../SonicTriton/python/customize.py | 42 +- .../SonicTriton/scripts/cmsTriton | 15 +- .../SonicTriton/src/RetryActionDiffServer.cc | 50 + .../src/RetryFallbackServerAction.cc | 54 + .../SonicTriton/src/TritonClient.cc | 1306 +++++++++-------- .../SonicTriton/src/TritonService.cc | 367 ++++- .../SonicTriton/test/BuildFile.xml | 22 +- .../test/DynamicModelLoadingProducer.cc | 83 ++ .../SonicTriton/test/RefCount.cc | 199 +++ .../SonicTriton/test/RetryActionDiffServer.cc | 71 + .../SonicTriton/test/tritonTest_cfg.py | 8 + 26 files changed, 1994 insertions(+), 795 deletions(-) create mode 100644 HeterogeneousCore/SonicCore/interface/RetryActionBase.h create mode 100644 HeterogeneousCore/SonicCore/plugins/BuildFile.xml create mode 100644 HeterogeneousCore/SonicCore/plugins/RetrySameServerAction.cc create mode 100644 HeterogeneousCore/SonicCore/src/RetryActionBase.cc create mode 100644 HeterogeneousCore/SonicTriton/interface/RetryActionDiffServer.h create mode 100644 HeterogeneousCore/SonicTriton/interface/RetryFallbackServerAction.h create mode 100644 HeterogeneousCore/SonicTriton/src/RetryActionDiffServer.cc create mode 100644 HeterogeneousCore/SonicTriton/src/RetryFallbackServerAction.cc create mode 100644 HeterogeneousCore/SonicTriton/test/DynamicModelLoadingProducer.cc create mode 100644 HeterogeneousCore/SonicTriton/test/RefCount.cc create mode 100644 HeterogeneousCore/SonicTriton/test/RetryActionDiffServer.cc diff --git a/HeterogeneousCore/SonicCore/BuildFile.xml b/HeterogeneousCore/SonicCore/BuildFile.xml index b0d5e2a08b98f..5208c91638f37 100644 --- a/HeterogeneousCore/SonicCore/BuildFile.xml +++ b/HeterogeneousCore/SonicCore/BuildFile.xml @@ -2,6 +2,7 @@ + diff --git a/HeterogeneousCore/SonicCore/interface/RetryActionBase.h b/HeterogeneousCore/SonicCore/interface/RetryActionBase.h new file mode 100644 index 0000000000000..ead6efd785dd2 --- /dev/null +++ b/HeterogeneousCore/SonicCore/interface/RetryActionBase.h @@ -0,0 +1,36 @@ +#ifndef HeterogeneousCore_SonicCore_RetryActionBase +#define HeterogeneousCore_SonicCore_RetryActionBase + +#include "FWCore/PluginManager/interface/PluginFactory.h" +#include "FWCore/ParameterSet/interface/ParameterSet.h" +#include "HeterogeneousCore/SonicCore/interface/SonicClientBase.h" +#include +#include + +// Base class for retry actions +class RetryActionBase { +public: + RetryActionBase(const edm::ParameterSet& conf, SonicClientBase* client); + virtual ~RetryActionBase() = default; + + bool shouldRetry() const { return shouldRetry_; } // Getter for shouldRetry_ + + virtual void retry() = 0; // Pure virtual function for execution logic + virtual void start() = 0; // Pure virtual function for execution logic for initialization + +protected: + void eval(); // interface for calling evaluate in client + void finish(bool success); // interface for calling finish directly in client + +protected: + SonicClientBase* client_; + bool shouldRetry_; // Flag to track if further retries should happen +}; + +// Define the factory for creating retry actions +using RetryActionFactory = + edmplugin::PluginFactory; + +#endif + +#define DEFINE_RETRY_ACTION(type) DEFINE_EDM_PLUGIN(RetryActionFactory, type, #type); diff --git a/HeterogeneousCore/SonicCore/interface/SonicClientBase.h b/HeterogeneousCore/SonicCore/interface/SonicClientBase.h index 47caaae8b2052..45a089701ed12 100644 --- a/HeterogeneousCore/SonicCore/interface/SonicClientBase.h +++ b/HeterogeneousCore/SonicCore/interface/SonicClientBase.h @@ -9,12 +9,15 @@ #include "HeterogeneousCore/SonicCore/interface/SonicDispatcherPseudoAsync.h" #include +#include #include #include #include enum class SonicMode { Sync = 1, Async = 2, PseudoAsync = 3 }; +class RetryActionBase; + class SonicClientBase { public: //constructor @@ -54,14 +57,23 @@ class SonicClientBase { SonicMode mode_; bool verbose_; std::unique_ptr dispatcher_; - unsigned allowedTries_, tries_; + unsigned totalTries_; std::optional holder_; + // Use a unique_ptr with a custom deleter to avoid incomplete type issues + struct RetryDeleter { + void operator()(RetryActionBase* ptr) const; + }; + + using RetryActionPtr = std::unique_ptr; + std::vector retryActions_; + //for logging/debugging std::string debugName_, clientName_, fullDebugName_; friend class SonicDispatcher; friend class SonicDispatcherPseudoAsync; + friend class RetryActionBase; }; #endif diff --git a/HeterogeneousCore/SonicCore/plugins/BuildFile.xml b/HeterogeneousCore/SonicCore/plugins/BuildFile.xml new file mode 100644 index 0000000000000..eaff0919e46bc --- /dev/null +++ b/HeterogeneousCore/SonicCore/plugins/BuildFile.xml @@ -0,0 +1,6 @@ + + + + + + diff --git a/HeterogeneousCore/SonicCore/plugins/RetrySameServerAction.cc b/HeterogeneousCore/SonicCore/plugins/RetrySameServerAction.cc new file mode 100644 index 0000000000000..09663ec8e4813 --- /dev/null +++ b/HeterogeneousCore/SonicCore/plugins/RetrySameServerAction.cc @@ -0,0 +1,28 @@ +#include "HeterogeneousCore/SonicCore/interface/RetryActionBase.h" +#include "HeterogeneousCore/SonicCore/interface/SonicClientBase.h" + +class RetrySameServerAction : public RetryActionBase { +public: + RetrySameServerAction(const edm::ParameterSet& pset, SonicClientBase* client) + : RetryActionBase(pset, client), allowedTries_(pset.getUntrackedParameter("allowedTries", 0)) {} + + void start() override { tries_ = 0; }; + +protected: + void retry() override; + +private: + unsigned allowedTries_, tries_; +}; + +void RetrySameServerAction::retry() { + ++tries_; + //if max retries has not been exceeded, call evaluate again + if (tries_ >= allowedTries_) { + shouldRetry_ = false; // Flip flag when max retries are reached + edm::LogInfo("RetrySameServerAction") << "Max retry attempts reached. No further retries."; + } + eval(); +} + +DEFINE_RETRY_ACTION(RetrySameServerAction) diff --git a/HeterogeneousCore/SonicCore/src/RetryActionBase.cc b/HeterogeneousCore/SonicCore/src/RetryActionBase.cc new file mode 100644 index 0000000000000..52b1fefb92a8a --- /dev/null +++ b/HeterogeneousCore/SonicCore/src/RetryActionBase.cc @@ -0,0 +1,15 @@ +#include "HeterogeneousCore/SonicCore/interface/RetryActionBase.h" + +// Constructor implementation +RetryActionBase::RetryActionBase(const edm::ParameterSet& conf, SonicClientBase* client) + : client_(client), shouldRetry_(true) { + if (client_ == nullptr) { + throw cms::Exception("RetryActionBase") << "client pointer cannot be null"; + } +} + +void RetryActionBase::eval() { client_->evaluate(); } + +void RetryActionBase::finish(bool success) { client_->finish(success); } + +EDM_REGISTER_PLUGINFACTORY(RetryActionFactory, "RetryActionFactory"); diff --git a/HeterogeneousCore/SonicCore/src/SonicClientBase.cc b/HeterogeneousCore/SonicCore/src/SonicClientBase.cc index 745c51f17aaf3..6ed10089e1bd3 100644 --- a/HeterogeneousCore/SonicCore/src/SonicClientBase.cc +++ b/HeterogeneousCore/SonicCore/src/SonicClientBase.cc @@ -1,18 +1,34 @@ #include "HeterogeneousCore/SonicCore/interface/SonicClientBase.h" +#include "HeterogeneousCore/SonicCore/interface/RetryActionBase.h" #include "FWCore/Utilities/interface/Exception.h" #include "FWCore/ParameterSet/interface/allowedValues.h" +// Custom deleter implementation +void SonicClientBase::RetryDeleter::operator()(RetryActionBase* ptr) const { delete ptr; } + SonicClientBase::SonicClientBase(const edm::ParameterSet& params, const std::string& debugName, const std::string& clientName) - : allowedTries_(params.getUntrackedParameter("allowedTries", 0)), - debugName_(debugName), - clientName_(clientName), - fullDebugName_(debugName_) { + : debugName_(debugName), clientName_(clientName), fullDebugName_(debugName_) { if (!clientName_.empty()) fullDebugName_ += ":" + clientName_; + const auto& retryPSetList = params.getParameter>("Retry"); std::string modeName(params.getParameter("mode")); + + for (const auto& retryPSet : retryPSetList) { + const std::string& actionType = retryPSet.getParameter("retryType"); + + auto retryAction = RetryActionFactory::get()->create(actionType, retryPSet, this); + if (retryAction) { + //Convert to RetryActionPtr Type from raw pointer of retryAction + retryActions_.emplace_back(RetryActionPtr(retryAction.release())); + } else { + throw cms::Exception("Configuration") + << "Unknown Retry type " << actionType << " for SonicClient: " << fullDebugName_; + } + } + if (modeName == "Sync") setMode(SonicMode::Sync); else if (modeName == "Async") @@ -40,24 +56,34 @@ void SonicClientBase::start(edm::WaitingTaskWithArenaHolder holder) { holder_ = std::move(holder); } -void SonicClientBase::start() { tries_ = 0; } +void SonicClientBase::start() { + totalTries_ = 0; + // initialize all actions + for (auto& action : retryActions_) { + action->start(); + } +} void SonicClientBase::finish(bool success, std::exception_ptr eptr) { //retries are only allowed if no exception was raised if (!success and !eptr) { - ++tries_; - //if max retries has not been exceeded, call evaluate again - if (tries_ < allowedTries_) { - evaluate(); - //avoid calling doneWaiting() twice - return; - } - //prepare an exception if exceeded - else { - edm::Exception ex(edm::errors::ExternalFailure); - ex << "SonicCallFailed: call failed after max " << tries_ << " tries"; - eptr = make_exception_ptr(ex); + ++totalTries_; + edm::LogInfo("SonicClientBase") << "finish: failed after total tries of " << totalTries_; + for (const auto& action : retryActions_) { + if (action->shouldRetry()) { + edm::LogInfo("SonicClientBase") << "Calling retry()"; + // retry() must trigger eval() or finish() + action->retry(); + // return because another finish() was already called inside client->evaluate() + return; + } } + //prepare an exception if no more retry actions left + edm::LogInfo("SonicClientBase") << "SonicCallFailed: call failed, no retry actions available after " << totalTries_ + << " tries."; + edm::Exception ex(edm::errors::ExternalFailure); + ex << "SonicCallFailed: call failed, no retry actions available after " << totalTries_ << " tries."; + eptr = make_exception_ptr(ex); } if (holder_) { holder_->doneWaiting(eptr); @@ -74,7 +100,20 @@ void SonicClientBase::fillBasePSetDescription(edm::ParameterSetDescription& desc //restrict allowed values desc.ifValue(edm::ParameterDescription("mode", "PseudoAsync", true), edm::allowedValues("Sync", "Async", "PseudoAsync")); - if (allowRetry) - desc.addUntracked("allowedTries", 0); + if (allowRetry) { + // Defines the structure of each entry in the VPSet + edm::ParameterSetDescription retryDesc; + retryDesc.add("retryType", "RetrySameServerAction"); + retryDesc.addUntracked("allowedTries", 0); + + // Define a default retry action + edm::ParameterSet defaultRetry; + defaultRetry.addParameter("retryType", "RetrySameServerAction"); + defaultRetry.addUntrackedParameter("allowedTries", 0); + + // Add the VPSet with the default retry action + desc.addVPSet("Retry", retryDesc, {defaultRetry}); + } + desc.add("sonicClientBase", desc); desc.addUntracked("verbose", false); } diff --git a/HeterogeneousCore/SonicCore/test/DummyClient.h b/HeterogeneousCore/SonicCore/test/DummyClient.h index ccef888ad9f7d..6504843926c0a 100644 --- a/HeterogeneousCore/SonicCore/test/DummyClient.h +++ b/HeterogeneousCore/SonicCore/test/DummyClient.h @@ -36,7 +36,7 @@ class DummyClient : public SonicClient { this->output_ = this->input_ * factor_; //simulate a failure - if (this->tries_ < fails_) + if (this->totalTries_ < fails_) this->finish(false); else this->finish(true); diff --git a/HeterogeneousCore/SonicCore/test/sonicTestAna_cfg.py b/HeterogeneousCore/SonicCore/test/sonicTestAna_cfg.py index 11c23c6cdfcc9..35cf42fa2b5ae 100644 --- a/HeterogeneousCore/SonicCore/test/sonicTestAna_cfg.py +++ b/HeterogeneousCore/SonicCore/test/sonicTestAna_cfg.py @@ -16,16 +16,27 @@ mode = cms.string("Sync"), factor = cms.int32(-1), wait = cms.int32(10), - allowedTries = cms.untracked.uint32(0), fails = cms.uint32(0), + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(0), + ) + ) ), ) process.dummySyncAnaRetry = process.dummySyncAna.clone( Client = dict( wait = 2, - allowedTries = 2, fails = 1, + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(2), + ) + ) + ) ) diff --git a/HeterogeneousCore/SonicCore/test/sonicTest_cfg.py b/HeterogeneousCore/SonicCore/test/sonicTest_cfg.py index 614297d86e3bb..bcbe820030440 100644 --- a/HeterogeneousCore/SonicCore/test/sonicTest_cfg.py +++ b/HeterogeneousCore/SonicCore/test/sonicTest_cfg.py @@ -17,15 +17,19 @@ process.options.numberOfThreads = 2 process.options.numberOfStreams = 0 - process.dummySync = _moduleClass(_moduleName, input = cms.int32(1), Client = cms.PSet( mode = cms.string("Sync"), factor = cms.int32(-1), wait = cms.int32(10), - allowedTries = cms.untracked.uint32(0), fails = cms.uint32(0), + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(0) + ) + ) ), ) @@ -35,8 +39,14 @@ mode = cms.string("PseudoAsync"), factor = cms.int32(2), wait = cms.int32(10), - allowedTries = cms.untracked.uint32(0), fails = cms.uint32(0), + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(0) + ) + ) + ), ) @@ -46,32 +56,53 @@ mode = cms.string("Async"), factor = cms.int32(5), wait = cms.int32(10), - allowedTries = cms.untracked.uint32(0), fails = cms.uint32(0), + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(0) + ) + ) ), ) process.dummySyncRetry = process.dummySync.clone( Client = dict( wait = 2, - allowedTries = 2, fails = 1, + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(2) + ) + ) + ) ) process.dummyPseudoAsyncRetry = process.dummyPseudoAsync.clone( Client = dict( wait = 2, - allowedTries = 2, fails = 1, + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(2) + ) + ) ) ) process.dummyAsyncRetry = process.dummyAsync.clone( Client = dict( wait = 2, - allowedTries = 2, fails = 1, + Retry = cms.VPSet( + cms.PSet( + allowedTries = cms.untracked.uint32(2), + retryType = cms.string('RetrySameServerAction') + ) + ) ) ) diff --git a/HeterogeneousCore/SonicTriton/BuildFile.xml b/HeterogeneousCore/SonicTriton/BuildFile.xml index b93d51e711e87..4af38d69d89e9 100644 --- a/HeterogeneousCore/SonicTriton/BuildFile.xml +++ b/HeterogeneousCore/SonicTriton/BuildFile.xml @@ -10,6 +10,7 @@ + - + diff --git a/HeterogeneousCore/SonicTriton/interface/RetryActionDiffServer.h b/HeterogeneousCore/SonicTriton/interface/RetryActionDiffServer.h new file mode 100644 index 0000000000000..4593201a473b3 --- /dev/null +++ b/HeterogeneousCore/SonicTriton/interface/RetryActionDiffServer.h @@ -0,0 +1,32 @@ +#ifndef HeterogeneousCore_SonicTriton_RetryActionDiffServer_h +#define HeterogeneousCore_SonicTriton_RetryActionDiffServer_h + +#include "HeterogeneousCore/SonicCore/interface/RetryActionBase.h" + +/** + * @class RetryActionDiffServer + * @brief A concrete implementation of RetryActionBase that attempts to retry an inference + * request on a different Triton server. + * + * This class provides a fallback mechanism. If an initial inference request fails + * (e.g., due to server unavailability or a model-specific error), this action will be + * triggered. It queries the central TritonService to select an alternative server (e.g., + * the fallback server when available) and instructs the TritonClient to reconnect to + * that server for the retry attempt. This action is designed for one-time use per + * inference call; after the retry attempt, it disables itself until the next `start()` + * call. + */ + +class RetryActionDiffServer : public RetryActionBase { +public: + RetryActionDiffServer(const edm::ParameterSet& conf, SonicClientBase* client); + ~RetryActionDiffServer() override = default; + + void retry() override; + void start() override; + +private: + unsigned tries_; +}; + +#endif diff --git a/HeterogeneousCore/SonicTriton/interface/RetryFallbackServerAction.h b/HeterogeneousCore/SonicTriton/interface/RetryFallbackServerAction.h new file mode 100644 index 0000000000000..8b72afcda08c8 --- /dev/null +++ b/HeterogeneousCore/SonicTriton/interface/RetryFallbackServerAction.h @@ -0,0 +1,18 @@ +#ifndef HeterogeneousCore_SonicTriton_RetryFallbackServerAction_h +#define HeterogeneousCore_SonicTriton_RetryFallbackServerAction_h + +#include "HeterogeneousCore/SonicCore/interface/RetryActionBase.h" + +class RetryFallbackServerAction : public RetryActionBase { +public: + RetryFallbackServerAction(const edm::ParameterSet& conf, SonicClientBase* client); + ~RetryFallbackServerAction() override = default; + + void retry() override; + void start() override; + +private: + unsigned tries_; +}; + +#endif diff --git a/HeterogeneousCore/SonicTriton/interface/TritonClient.h b/HeterogeneousCore/SonicTriton/interface/TritonClient.h index 44118d43d09f1..74d6649f6c5e8 100644 --- a/HeterogeneousCore/SonicTriton/interface/TritonClient.h +++ b/HeterogeneousCore/SonicTriton/interface/TritonClient.h @@ -1,104 +1,118 @@ -#ifndef HeterogeneousCore_SonicTriton_TritonClient -#define HeterogeneousCore_SonicTriton_TritonClient - -#include "FWCore/ParameterSet/interface/ParameterSet.h" -#include "FWCore/ParameterSet/interface/ParameterSetDescription.h" -#include "FWCore/ServiceRegistry/interface/ServiceToken.h" -#include "HeterogeneousCore/SonicCore/interface/SonicClient.h" -#include "HeterogeneousCore/SonicTriton/interface/TritonData.h" -#include "HeterogeneousCore/SonicTriton/interface/TritonService.h" - -#include -#include -#include -#include -#include - -#include "grpc_client.h" -#include "grpc_service.pb.h" - -enum class TritonBatchMode { Rectangular = 1, Ragged = 2 }; - -class TritonClient : public SonicClient { -public: - struct ServerSideStats { - uint64_t inference_count_; - uint64_t execution_count_; - uint64_t success_count_; - uint64_t cumm_time_ns_; - uint64_t queue_time_ns_; - uint64_t compute_input_time_ns_; - uint64_t compute_infer_time_ns_; - uint64_t compute_output_time_ns_; - }; - - //constructor - TritonClient(const edm::ParameterSet& params, const std::string& debugName); - - //destructor - ~TritonClient() override; - - //accessors - unsigned batchSize() const; - TritonBatchMode batchMode() const { return batchMode_; } - bool verbose() const { return verbose_; } - bool useSharedMemory() const { return useSharedMemory_; } - void setUseSharedMemory(bool useShm) { useSharedMemory_ = useShm; } - bool setBatchSize(unsigned bsize); - void setBatchMode(TritonBatchMode batchMode); - void resetBatchMode(); - void reset() override; - TritonServerType serverType() const { return serverType_; } - bool isLocal() const { return isLocal_; } - const TritonService* service() const; - const TritonService* localService() const; - - //for fillDescriptions - static void fillPSetDescription(edm::ParameterSetDescription& iDesc); - -protected: - //helpers - bool noOuterDim() const { return noOuterDim_; } - unsigned outerDim() const { return outerDim_; } - unsigned nEntries() const; - void getResults(const std::vector>& results); - void evaluate() override; - template - bool handle_exception(F&& call); - - void reportServerSideStats(const ServerSideStats& stats) const; - ServerSideStats summarizeServerStats(const inference::ModelStatistics& start_status, - const inference::ModelStatistics& end_status) const; - - inference::ModelStatistics getServerSideStatus() const; - - //members - unsigned maxOuterDim_; - unsigned outerDim_; - bool noOuterDim_; - unsigned nEntries_; - TritonBatchMode batchMode_; - bool manualBatchMode_; - bool verbose_; - bool useSharedMemory_; - TritonServerType serverType_; - bool isLocal_; - grpc_compression_algorithm compressionAlgo_; - triton::client::Headers headers_; - - std::unique_ptr client_; - //stores timeout, model name and version - std::vector options_; - edm::ServiceToken token_; - -private: - friend TritonInputData; - friend TritonOutputData; - - //private accessors only used by data - auto client() { return client_.get(); } - void addEntry(unsigned entry); - void resizeEntries(unsigned entry); -}; - -#endif +#ifndef HeterogeneousCore_SonicTriton_TritonClient +#define HeterogeneousCore_SonicTriton_TritonClient + +#include "FWCore/ParameterSet/interface/ParameterSet.h" +#include "FWCore/ParameterSet/interface/ParameterSetDescription.h" +#include "FWCore/Concurrency/interface/WaitingTaskWithArenaHolder.h" +#include "FWCore/ServiceRegistry/interface/ServiceToken.h" +#include "HeterogeneousCore/SonicCore/interface/SonicClient.h" +#include "HeterogeneousCore/SonicTriton/interface/TritonData.h" +#include "HeterogeneousCore/SonicTriton/interface/TritonService.h" + +#include +#include +#include +#include +#include +#include + +#include "grpc_client.h" +#include "grpc_service.pb.h" + +enum class TritonBatchMode { Rectangular = 1, Ragged = 2 }; + +class TritonClient : public SonicClient { +public: + struct ServerSideStats { + uint64_t inference_count_; + uint64_t execution_count_; + uint64_t success_count_; + uint64_t cumm_time_ns_; + uint64_t queue_time_ns_; + uint64_t compute_input_time_ns_; + uint64_t compute_infer_time_ns_; + uint64_t compute_output_time_ns_; + }; + + //constructor + TritonClient(const edm::ParameterSet& params, const std::string& debugName); + + //destructor + ~TritonClient() override; + + //accessors + unsigned batchSize() const; + TritonBatchMode batchMode() const { return batchMode_; } + bool verbose() const { return verbose_; } + bool useSharedMemory() const { return useSharedMemory_; } + void setUseSharedMemory(bool useShm) { useSharedMemory_ = useShm; } + bool setBatchSize(unsigned bsize); + void setBatchMode(TritonBatchMode batchMode); + void resetBatchMode(); + void reset() override; + TritonServerType serverType() const { return serverType_; } + bool isLocal() const { return isLocal_; } + const TritonService* service() const; + const TritonService* localService() const; + std::string modelName() const { return options_[0].model_name_; } + std::string serverName() const { return serverName_; } + virtual void connectToServer(const std::string& url); + virtual void updateServer(const std::string& serverName); + virtual void switchToFallback(); + + //for fillDescriptions + static void fillPSetDescription(edm::ParameterSetDescription& iDesc); + +protected: + // Protected default constructor for unit testing (no framework services) + TritonClient(); + + //helpers + bool noOuterDim() const { return noOuterDim_; } + unsigned outerDim() const { return outerDim_; } + unsigned nEntries() const; + void getResults(const std::vector>& results); + void evaluate() override; + template + bool handle_exception(F&& call); + template + bool handle_exception_holder(edm::WaitingTaskWithArenaHolder& fh, F&& call); + + void reportServerSideStats(const ServerSideStats& stats) const; + ServerSideStats summarizeServerStats(const inference::ModelStatistics& start_status, + const inference::ModelStatistics& end_status) const; + + inference::ModelStatistics getServerSideStatus() const; + + //members + unsigned maxOuterDim_; + unsigned outerDim_; + bool noOuterDim_; + unsigned nEntries_; + TritonBatchMode batchMode_; + bool manualBatchMode_; + bool verbose_; + bool useSharedMemory_; + TritonServerType serverType_; + bool isLocal_; + std::string serverName_; + grpc_compression_algorithm compressionAlgo_; + triton::client::Headers headers_; + + std::unique_ptr client_; + //stores timeout, model name and version + std::vector options_; + edm::ServiceToken token_; + std::atomic inferSuccess_{true}; + +private: + friend TritonInputData; + friend TritonOutputData; + + //private accessors only used by data + auto client() { return client_.get(); } + void addEntry(unsigned entry); + void resizeEntries(unsigned entry); +}; + +#endif diff --git a/HeterogeneousCore/SonicTriton/interface/TritonService.h b/HeterogeneousCore/SonicTriton/interface/TritonService.h index 8ac7a915f8d6d..9be006dfab615 100644 --- a/HeterogeneousCore/SonicTriton/interface/TritonService.h +++ b/HeterogeneousCore/SonicTriton/interface/TritonService.h @@ -3,6 +3,7 @@ #include "FWCore/ParameterSet/interface/ParameterSet.h" #include "FWCore/Utilities/interface/GlobalIdentifier.h" +#include "oneapi/tbb/concurrent_hash_map.h" #include #include @@ -11,6 +12,8 @@ #include #include #include +#include +#include #include "grpc_client.h" @@ -90,18 +93,28 @@ class TritonService { static const std::string fallbackAddress; static const std::string siteconfName; }; + //Dynamic quantities of servers + struct ServerHealth { + bool live{false}; + bool ready{false}; + + uint64_t inferenceCount{0}; + uint64_t failureCount{0}; + double avgQueueTimeMs{0.0}; + double avgInferTimeMs{0.0}; + }; struct Model { Model(const std::string& path_ = "") : path(path_) {} - //members std::string path; std::unordered_set servers; std::unordered_set modules; + int refCount{0}; // for dynamic loading on fallback server + bool isLoaded() const { return refCount > 0; } }; struct Module { //currently assumes that a module can only have one associated model Module(const std::string& model_) : model(model_) {} - //members std::string model; }; @@ -111,12 +124,40 @@ class TritonService { //accessors void addModel(const std::string& modelName, const std::string& path); - Server serverInfo(const std::string& model, const std::string& preferred = "") const; + + const std::string* resolveServerName(const std::string& model, const std::string& preferred = "") const; + const std::pair& resolveServer(const std::string& model, + const std::string& preferred = "") const; + std::vector unassignedModels() const; + + // update health stats of all servers + void updateServerHealth(const std::string& modelName = "") const; + + // return the best server for retry, ignore the current server + std::optional getBestServer(const std::string& modelName, const std::string& IgnoreServer = "") const; + + // helper functions to get server statistics? + // - getServerSideStatus() + // - updateServerStatus() + // - loop over servers_ get statistics + // - getBestServer(model) + // - call updateServerStatus() + // - loop over servers_ get their statistics, compute metric, return server name + const std::string& pid() const { return pid_; } void notifyCallStatus(bool status) const; static void fillDescriptions(edm::ConfigurationDescriptions& descriptions); + // Dynamic model loading/unloading - only supported for the fallback server + // The fallback server must be started with explicit model control mode + // (--model-control-mode explicit) for these functions to work + bool loadModel(const std::string& modelName); + bool unloadModel(const std::string& modelName); + // Start the fallback server if enabled and not already running (idempotent) + void startFallbackServer(); + bool fallbackStarted() const { return startedFallback_; } + private: void preallocate(edm::service::SystemBounds const&); void preModuleConstruction(edm::ModuleDescription const&); @@ -128,6 +169,9 @@ class TritonService { //helper template void printFallbackServerLog() const; + // Internal helpers that operate on Model directly (caller holds lock) + bool loadModel(const std::string& modelName, Model& model); + bool unloadModel(const std::string& modelName, Model& model); bool verbose_; FallbackOpts fallbackOpts_; @@ -136,12 +180,18 @@ class TritonService { bool startedFallback_; mutable std::atomic callFails_; std::string pid_; - std::unordered_map unservedModels_; //this represents a many:many:many map std::unordered_map servers_; + //server health needs concurrent-safe edits + tbb::concurrent_hash_map serversHealth_; std::unordered_map models_; std::unordered_map modules_; int numberOfThreads_; + + //Dynamic model loading and unloading (fallback server only) + std::mutex modelLoadMutex_; + // Model names currently loaded on the fallback server + std::unordered_set fallbackLoadedModels_; }; #endif diff --git a/HeterogeneousCore/SonicTriton/python/customize.py b/HeterogeneousCore/SonicTriton/python/customize.py index 9873a1a498dd7..e633d169e919e 100644 --- a/HeterogeneousCore/SonicTriton/python/customize.py +++ b/HeterogeneousCore/SonicTriton/python/customize.py @@ -13,12 +13,12 @@ def getParser(): from argparse import ArgumentParser, ArgumentDefaultsHelpFormatter parser = ArgumentParser(formatter_class=ArgumentDefaultsHelpFormatter) parser.add_argument("--maxEvents", default=-1, type=int, help="Number of events to process (-1 for all)") - parser.add_argument("--serverName", default="default", type=str, help="name for server (used internally)") - parser.add_argument("--address", default="", type=str, help="server address") - parser.add_argument("--port", default=8001, type=int, help="server port") + parser.add_argument("--address", nargs=3, action="append", metavar=("NAME", "HOST", "PORT"), + dest="addresses", default=[], + help="Triton server entry: name host port (repeatable, e.g. --address server1 0.0.0.0 8011)") parser.add_argument("--timeout", default=30, type=int, help="timeout for requests") parser.add_argument("--timeoutUnit", default="seconds", type=str, help="unit for timeout") - parser.add_argument("--params", default="", type=str, help="json file containing server address/port") + parser.add_argument("--params", default="", type=str, help="json file containing server address/port(single-server)") parser.add_argument("--threads", default=1, type=int, help="number of threads") parser.add_argument("--streams", default=0, type=int, help="number of streams") parser.add_argument("--verbose", default=False, action="store_true", help="enable all verbose output") @@ -36,12 +36,27 @@ def getParser(): parser.add_argument("--imageName", default="", type=str, help="container image name for fallback server") parser.add_argument("--sandboxDir", default="", type=str, help="apptainer sandbox directory") parser.add_argument("--tempDir", default="", type=str, help="temp directory for fallback server") + parser.add_argument("--retryAction", default="same", type=str, choices=["same","diff"], help="retry policy: same server or different server") return parser def getOptions(parser, verbose=False): options = parser.parse_args() + # Legacy --params support: loads a single server and appends it to options.addresses + if len(options.params) > 0: + with open(options.params, 'r') as pfile: + pdict = json.load(pfile) + name = pdict.get("name", "default") + host = pdict["address"] + port = str(int(pdict["port"])) + options.addresses.append([name, host, port]) + if verbose: + print("server (from params) = {}:{} [{}]".format(host, port, name)) + + if verbose: + for name, host, port in options.addresses: + print("server = {}:{} [{}]".format(host, port, name)) if len(options.params)>0: with open(options.params,'r') as pfile: pdict = json.load(pfile) @@ -76,13 +91,13 @@ def applyOptions(process, options, applyToModules=False): process.TritonService.fallback.device = options.device if len(options.fallbackName)>0: process.TritonService.fallback.instanceBaseName = options.fallbackName - if len(options.address)>0: + for name, host, port in options.addresses: process.TritonService.servers.append( dict( - name = options.serverName, - address = options.address, - port = options.port, - useSsl = options.ssl, + name = name, + address = host, + port = int(port), + useSsl = options.ssl, ) ) @@ -92,12 +107,19 @@ def applyOptions(process, options, applyToModules=False): return process def getClientOptions(options): + action = cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(options.tries)) + if options.retryAction != 'same': + action.retryType = cms.string('RetryActionDiffServer') + + fallback = cms.PSet(retryType = cms.string('RetryFallbackServerAction')) return dict( compression = cms.untracked.string(options.compression), useSharedMemory = cms.untracked.bool(not options.noShm), timeout = cms.untracked.uint32(options.timeout), timeoutUnit = cms.untracked.string(options.timeoutUnit), - allowedTries = cms.untracked.uint32(options.tries), + Retry = cms.VPSet(action,fallback) ) def applyClientOptions(client, options): diff --git a/HeterogeneousCore/SonicTriton/scripts/cmsTriton b/HeterogeneousCore/SonicTriton/scripts/cmsTriton index 01932254f9f1d..c3d7a9fc01d87 100755 --- a/HeterogeneousCore/SonicTriton/scripts/cmsTriton +++ b/HeterogeneousCore/SonicTriton/scripts/cmsTriton @@ -31,6 +31,7 @@ fi DEVICE=auto THREADCONTROL="" GLOBALCONNECT= +MODELCONTROL="explicit" get_sandbox(){ if [ -z "$SANDBOX" ]; then @@ -55,6 +56,7 @@ usage() { $ECHO "-G \t accept connections globally (default: only accept connections from localhost)" $ECHO "-i [name] \t server image name (default: ${IMAGE})" $ECHO "-I [num] \t number of model instances (default: ${INSTANCES} -> means no local editing of config files)" + $ECHO "-L \t disable explicit model control mode (enabled by default for dynamic model loading)" $ECHO "-M [dir] \t model repository (can be given more than once)" $ECHO "-m [dir] \t specific model directory (can be given more than once)" $ECHO "-n [name] \t name of container instance, also used for default hidden temporary dir (default: ${SERVER})" @@ -80,7 +82,7 @@ if [ -e /run/shm ]; then SHM=/run/shm fi -while getopts "cC:Dd:fg:i:I:M:m:n:P:p:r:s:t:vw:h" opt; do +while getopts "cC:Dd:fg:i:I:LM:m:n:P:p:r:s:t:vw:h" opt; do case "$opt" in c) CLEANUP="" ;; @@ -100,6 +102,8 @@ while getopts "cC:Dd:fg:i:I:M:m:n:P:p:r:s:t:vw:h" opt; do ;; I) INSTANCES="$OPTARG" ;; + L) MODELCONTROL="none" + ;; M) REPOS+=("$OPTARG") ;; m) MODELS+=("$OPTARG") @@ -250,6 +254,9 @@ start_docker(){ # mount all model repositories MOUNTARGS="" REPOARGS="" + if [ -n "$MODELCONTROL" ]; then + REPOARGS="--model-control-mode=${MODELCONTROL}" + fi for REPO in ${REPOS[@]}; do MOUNTARGS="$MOUNTARGS -v$REPO:$REPO" REPOARGS="$REPOARGS --model-repository=${REPO}" @@ -276,6 +283,9 @@ start_podman(){ # mount all model repositories MOUNTARGS="" REPOARGS="" + if [ -n "$MODELCONTROL" ]; then + REPOARGS="--model-control-mode=${MODELCONTROL}" + fi for REPO in ${REPOS[@]}; do MOUNTARGS="$MOUNTARGS --volume $REPO:$REPO" REPOARGS="$REPOARGS --model-repository=${REPO}" @@ -302,6 +312,9 @@ start_apptainer(){ # mount all model repositories MOUNTARGS="" REPOARGS="" + if [ -n "$MODELCONTROL" ]; then + REPOARGS="--model-control-mode=${MODELCONTROL}" + fi for REPO in ${REPOS[@]}; do MOUNTARGS="$MOUNTARGS -B $REPO" REPOARGS="$REPOARGS --model-repository=${REPO}" diff --git a/HeterogeneousCore/SonicTriton/src/RetryActionDiffServer.cc b/HeterogeneousCore/SonicTriton/src/RetryActionDiffServer.cc new file mode 100644 index 0000000000000..8a096f1bfd672 --- /dev/null +++ b/HeterogeneousCore/SonicTriton/src/RetryActionDiffServer.cc @@ -0,0 +1,50 @@ +#include "HeterogeneousCore/SonicTriton/interface/RetryActionDiffServer.h" +#include "HeterogeneousCore/SonicTriton/interface/TritonClient.h" +#include "HeterogeneousCore/SonicTriton/interface/TritonService.h" +#include "FWCore/MessageLogger/interface/MessageLogger.h" +#include "FWCore/ServiceRegistry/interface/Service.h" + +RetryActionDiffServer::RetryActionDiffServer(const edm::ParameterSet& conf, SonicClientBase* client) + : RetryActionBase(conf, client) {} + +void RetryActionDiffServer::start() { + this->shouldRetry_ = true; + tries_ = 0; +} + +void RetryActionDiffServer::retry() { + ++tries_; + if (tries_ >= 1) { + shouldRetry_ = false; // Flip flag when max retries are reached. Allow 1 try for now. + edm::LogInfo("RetryDiffServerAction") << "Max retry attempts reached. No further retries."; + } + try { + auto* tritonClient = static_cast(client_); + edm::LogInfo("RetryActionDiffServer") << "Asking for a different server from TritonService"; + auto ts = tritonClient->service(); + + // First, try to find another remote server + auto bestServerName = ts->getBestServer(tritonClient->modelName(), tritonClient->serverName()); + + if (bestServerName) { + edm::LogInfo("RetryActionDiffServer") << "Got best server from service "; + tritonClient->updateServer(*bestServerName); + edm::LogInfo("RetryActionDiffServer") << "eval() with new server"; + eval(); + return; + } else { + edm::LogWarning("RetryActionDiffServer") + << "No alternative server found for model " << tritonClient->modelName() << ". Now call client->finish()"; + finish(false); + return; + } + } catch (TritonException& e) { + e.convertToWarning(); + } catch (std::exception& e) { + edm::LogError("RetryActionDiffServer") << "Failed to retry with alternative server: " << e.what(); + } catch (...) { + edm::LogError("RetryActionDiffServer: UnknownFailure") << "An unknown exception was thrown"; + } +} + +DEFINE_RETRY_ACTION(RetryActionDiffServer); diff --git a/HeterogeneousCore/SonicTriton/src/RetryFallbackServerAction.cc b/HeterogeneousCore/SonicTriton/src/RetryFallbackServerAction.cc new file mode 100644 index 0000000000000..823d6ce4eca2b --- /dev/null +++ b/HeterogeneousCore/SonicTriton/src/RetryFallbackServerAction.cc @@ -0,0 +1,54 @@ +// RetryFallbackServerAction: last-resort retry action for TritonClient. +// +// When all other retry actions have been exhausted, this action loads the +// client's model onto the fallback (local) Triton server and re-runs +// inference there. It fires at most once per inference call. +// +// Usage add to the Retry VPSet *after* all other retry actions: +// cms.PSet(retryType = cms.string("RetryFallbackServerAction")) +// +// Requirements: +// - TritonService fallback must be enabled in the job configuration. +// - The model must have a modelConfigPath / repository path known to +// TritonService so it can be loaded dynamically. + +#include "HeterogeneousCore/SonicTriton/interface/RetryFallbackServerAction.h" +#include "HeterogeneousCore/SonicTriton/interface/TritonClient.h" +#include "HeterogeneousCore/SonicTriton/interface/TritonService.h" +#include "FWCore/MessageLogger/interface/MessageLogger.h" +#include "FWCore/ParameterSet/interface/ParameterSet.h" +#include "FWCore/Utilities/interface/Exception.h" + +RetryFallbackServerAction::RetryFallbackServerAction(const edm::ParameterSet& conf, SonicClientBase* client) + : RetryActionBase(conf, client) {} + +void RetryFallbackServerAction::start() { + this->shouldRetry_ = true; + tries_ = 0; +} + +void RetryFallbackServerAction::retry() { + // Allow only one fallback attempt per inference call. + shouldRetry_ = false; + + auto* tc = dynamic_cast(client_); + if (!tc) { + // Should never happen in a correctly configured job. + edm::LogWarning("RetryFallbackServerAction") + << "client_ is not a TritonClient — cannot redirect to fallback server"; + finish(false); + return; + } + + CMS_SA_ALLOW try { + // Start the fallback server (idempotent), load the model, and point + // the client's gRPC connection at the fallback URL. + tc->switchToFallback(); + // Re-run the inference on the fallback server. + eval(); + } catch (...) { + // Non-retryable: propagate the exception so the job fails cleanly. + finish(false); + } +} +DEFINE_RETRY_ACTION(RetryFallbackServerAction); diff --git a/HeterogeneousCore/SonicTriton/src/TritonClient.cc b/HeterogeneousCore/SonicTriton/src/TritonClient.cc index 05fb59ce9b0c0..bc223440a68b4 100644 --- a/HeterogeneousCore/SonicTriton/src/TritonClient.cc +++ b/HeterogeneousCore/SonicTriton/src/TritonClient.cc @@ -1,601 +1,705 @@ -#include "FWCore/MessageLogger/interface/MessageLogger.h" -#include "FWCore/ParameterSet/interface/FileInPath.h" -#include "FWCore/ParameterSet/interface/allowedValues.h" -#include "FWCore/ServiceRegistry/interface/Service.h" -#include "FWCore/Utilities/interface/Exception.h" -#include "HeterogeneousCore/SonicTriton/interface/TritonClient.h" -#include "HeterogeneousCore/SonicTriton/interface/TritonException.h" -#include "HeterogeneousCore/SonicTriton/interface/TritonService.h" -#include "HeterogeneousCore/SonicTriton/interface/triton_utils.h" - -#include "grpc_client.h" -#include "grpc_service.pb.h" -#include "model_config.pb.h" - -#include "google/protobuf/text_format.h" -#include "google/protobuf/io/zero_copy_stream_impl.h" - -#include -#include -#include -#include -#include -#include -#include -#include -#include - -namespace tc = triton::client; - -namespace { - grpc_compression_algorithm getCompressionAlgo(const std::string& name) { - if (name.empty() or name.compare("none") == 0) - return grpc_compression_algorithm::GRPC_COMPRESS_NONE; - else if (name.compare("deflate") == 0) - return grpc_compression_algorithm::GRPC_COMPRESS_DEFLATE; - else if (name.compare("gzip") == 0) - return grpc_compression_algorithm::GRPC_COMPRESS_GZIP; - else - throw cms::Exception("GrpcCompression") - << "Unknown compression algorithm requested: " << name << " (choices: none, deflate, gzip)"; - } - - std::vector> convertToShared(const std::vector& tmp) { - std::vector> results; - results.reserve(tmp.size()); - std::transform(tmp.begin(), tmp.end(), std::back_inserter(results), [](tc::InferResult* ptr) { - return std::shared_ptr(ptr); - }); - return results; - } -} // namespace - -//based on https://github.com/triton-inference-server/server/blob/v2.3.0/src/clients/c++/examples/simple_grpc_async_infer_client.cc -//and https://github.com/triton-inference-server/server/blob/v2.3.0/src/clients/c++/perf_client/perf_client.cc - -TritonClient::TritonClient(const edm::ParameterSet& params, const std::string& debugName) - : SonicClient(params, debugName, "TritonClient"), - batchMode_(TritonBatchMode::Rectangular), - manualBatchMode_(false), - verbose_(params.getUntrackedParameter("verbose")), - useSharedMemory_(params.getUntrackedParameter("useSharedMemory")), - compressionAlgo_(getCompressionAlgo(params.getUntrackedParameter("compression"))) { - options_.emplace_back(params.getParameter("modelName")); - //get appropriate server for this model - edm::Service ts; - - // We save the token to be able to notify the service in case of an exception in the evaluate method. - // The evaluate method can be called outside the frameworks TBB threadpool in the case of a retry. In - // this case the context is not setup to access the service registry, we need the service token to - // create the context. - token_ = edm::ServiceRegistry::instance().presentToken(); - - const auto& server = - ts->serverInfo(options_[0].model_name_, params.getUntrackedParameter("preferredServer")); - serverType_ = server.type; - edm::LogInfo("TritonDiscovery") << debugName_ << " assigned server: " << server.url; - //enforce sync mode for fallback CPU server to avoid contention - //todo: could enforce async mode otherwise (unless mode was specified by user?) - if (serverType_ == TritonServerType::LocalCPU) - setMode(SonicMode::Sync); - isLocal_ = serverType_ == TritonServerType::LocalCPU or serverType_ == TritonServerType::LocalGPU; - - //connect to the server - TRITON_THROW_IF_ERROR( - tc::InferenceServerGrpcClient::Create(&client_, server.url, false, server.useSsl, server.sslOptions), - "TritonClient(): unable to create inference context", - localService()); - - //set options - options_[0].model_version_ = params.getParameter("modelVersion"); - options_[0].client_timeout_ = params.getUntrackedParameter("timeout"); - //convert to microseconds - const auto& timeoutUnit = params.getUntrackedParameter("timeoutUnit"); - unsigned conversion = 1; - if (timeoutUnit == "seconds") - conversion = 1e6; - else if (timeoutUnit == "milliseconds") - conversion = 1e3; - else if (timeoutUnit == "microseconds") - conversion = 1; - else - throw cms::Exception("Configuration") << "Unknown timeout unit: " << timeoutUnit; - options_[0].client_timeout_ *= conversion; - - //get fixed parameters from local config - inference::ModelConfig localModelConfig; - { - const std::string localModelConfigPath(params.getParameter("modelConfigPath").fullPath()); - int fileDescriptor = open(localModelConfigPath.c_str(), O_RDONLY); - if (fileDescriptor < 0) - throw TritonException("LocalFailure") - << "TritonClient(): unable to open local model config: " << localModelConfigPath; - google::protobuf::io::FileInputStream localModelConfigInput(fileDescriptor); - localModelConfigInput.SetCloseOnDelete(true); - if (!google::protobuf::TextFormat::Parse(&localModelConfigInput, &localModelConfig)) - throw TritonException("LocalFailure") - << "TritonClient(): unable to parse local model config: " << localModelConfigPath; - } - - //check batch size limitations (after i/o setup) - //triton uses max batch size = 0 to denote a model that does not support native batching (using the outer dimension) - //but for models that do support batching (native or otherwise), a given event may set batch size 0 to indicate no valid input is present - //so set the local max to 1 and keep track of "no outer dim" case - maxOuterDim_ = localModelConfig.max_batch_size(); - noOuterDim_ = maxOuterDim_ == 0; - maxOuterDim_ = std::max(1u, maxOuterDim_); - //propagate batch size - setBatchSize(1); - - //compare model checksums to remote config to enforce versioning - inference::ModelConfigResponse modelConfigResponse; - TRITON_THROW_IF_ERROR(client_->ModelConfig(&modelConfigResponse, options_[0].model_name_, options_[0].model_version_), - "TritonClient(): unable to get model config", - localService()); - inference::ModelConfig remoteModelConfig(modelConfigResponse.config()); - - std::map> checksums; - size_t fileCounter = 0; - for (const auto& modelConfig : {localModelConfig, remoteModelConfig}) { - const auto& agents = modelConfig.model_repository_agents().agents(); - auto agent = std::find_if(agents.begin(), agents.end(), [](auto const& a) { return a.name() == "checksum"; }); - if (agent != agents.end()) { - const auto& params = agent->parameters(); - for (const auto& [key, val] : params) { - // only check the requested version - if (key.compare(0, options_[0].model_version_.size() + 1, options_[0].model_version_ + "/") == 0) - checksums[key][fileCounter] = val; - } - } - ++fileCounter; - } - std::vector incorrect; - for (const auto& [key, val] : checksums) { - if (checksums[key][0] != checksums[key][1]) - incorrect.push_back(key); - } - if (!incorrect.empty()) - throw TritonException("ModelVersioning") << "The following files have incorrect checksums on the remote server: " - << triton_utils::printColl(incorrect, ", "); - - //get model info - inference::ModelMetadataResponse modelMetadata; - TRITON_THROW_IF_ERROR(client_->ModelMetadata(&modelMetadata, options_[0].model_name_, options_[0].model_version_), - "TritonClient(): unable to get model metadata", - localService()); - - //get input and output (which know their sizes) - const auto& nicInputs = modelMetadata.inputs(); - const auto& nicOutputs = modelMetadata.outputs(); - - //report all model errors at once - std::stringstream msg; - std::string msg_str; - - //currently no use case is foreseen for a model with zero inputs or outputs - if (nicInputs.empty()) - msg << "Model on server appears malformed (zero inputs)\n"; - - if (nicOutputs.empty()) - msg << "Model on server appears malformed (zero outputs)\n"; - - //stop if errors - msg_str = msg.str(); - if (!msg_str.empty()) - throw cms::Exception("ModelErrors") << msg_str; - - //setup input map - std::stringstream io_msg; - if (verbose_) - io_msg << "Model inputs: " - << "\n"; - for (const auto& nicInput : nicInputs) { - const auto& iname = nicInput.name(); - auto [curr_itr, success] = input_.emplace(std::piecewise_construct, - std::forward_as_tuple(iname), - std::forward_as_tuple(iname, nicInput, this, ts->pid())); - auto& curr_input = curr_itr->second; - if (verbose_) { - io_msg << " " << iname << " (" << curr_input.dname() << ", " << curr_input.byteSize() - << " b) : " << triton_utils::printColl(curr_input.shape()) << "\n"; - } - } - - //allow selecting only some outputs from server - const auto& v_outputs = params.getUntrackedParameter>("outputs"); - std::unordered_set s_outputs(v_outputs.begin(), v_outputs.end()); - - //setup output map - if (verbose_) - io_msg << "Model outputs: " - << "\n"; - for (const auto& nicOutput : nicOutputs) { - const auto& oname = nicOutput.name(); - if (!s_outputs.empty() and s_outputs.find(oname) == s_outputs.end()) - continue; - auto [curr_itr, success] = output_.emplace(std::piecewise_construct, - std::forward_as_tuple(oname), - std::forward_as_tuple(oname, nicOutput, this, ts->pid())); - auto& curr_output = curr_itr->second; - if (verbose_) { - io_msg << " " << oname << " (" << curr_output.dname() << ", " << curr_output.byteSize() - << " b) : " << triton_utils::printColl(curr_output.shape()) << "\n"; - } - if (!s_outputs.empty()) - s_outputs.erase(oname); - } - - //check if any requested outputs were not available - if (!s_outputs.empty()) - throw cms::Exception("MissingOutput") - << "Some requested outputs were not available on the server: " << triton_utils::printColl(s_outputs); - - //print model info - std::stringstream model_msg; - if (verbose_) { - model_msg << "Model name: " << options_[0].model_name_ << "\n" - << "Model version: " << options_[0].model_version_ << "\n" - << "Model max outer dim: " << (noOuterDim_ ? 0 : maxOuterDim_) << "\n"; - edm::LogInfo(fullDebugName_) << model_msg.str() << io_msg.str(); - } -} - -TritonClient::~TritonClient() { - //by default: members of this class destroyed before members of base class - //in shared memory case, TritonMemResource (member of TritonData) unregisters from client_ in its destructor - //but input/output objects are member of base class, so destroyed after client_ (member of this class) - //therefore, clear the maps here - input_.clear(); - output_.clear(); -} - -void TritonClient::setBatchMode(TritonBatchMode batchMode) { - unsigned oldBatchSize = batchSize(); - batchMode_ = batchMode; - manualBatchMode_ = true; - //this allows calling setBatchSize() and setBatchMode() in either order consistently to change back and forth - //includes handling of change from ragged to rectangular if multiple entries already created - setBatchSize(oldBatchSize); -} - -void TritonClient::resetBatchMode() { - batchMode_ = TritonBatchMode::Rectangular; - manualBatchMode_ = false; -} - -unsigned TritonClient::nEntries() const { return !input_.empty() ? input_.begin()->second.entries_.size() : 0; } - -unsigned TritonClient::batchSize() const { return batchMode_ == TritonBatchMode::Rectangular ? outerDim_ : nEntries(); } - -bool TritonClient::setBatchSize(unsigned bsize) { - if (batchMode_ == TritonBatchMode::Rectangular) { - if (bsize > maxOuterDim_) { - throw TritonException("LocalFailure") - << "Requested batch size " << bsize << " exceeds server-specified max batch size " << maxOuterDim_ << "."; - return false; - } else { - outerDim_ = bsize; - //take min to allow resizing to 0 - resizeEntries(std::min(outerDim_, 1u)); - return true; - } - } else { - resizeEntries(bsize); - outerDim_ = 1; - return true; - } -} - -void TritonClient::resizeEntries(unsigned entry) { - if (entry > nEntries()) - //addEntry(entry) extends the vector to size entry+1 - addEntry(entry - 1); - else if (entry < nEntries()) { - for (auto& element : input_) { - element.second.entries_.resize(entry); - } - for (auto& element : output_) { - element.second.entries_.resize(entry); - } - } -} - -void TritonClient::addEntry(unsigned entry) { - for (auto& element : input_) { - element.second.addEntryImpl(entry); - } - for (auto& element : output_) { - element.second.addEntryImpl(entry); - } - if (entry > 0) { - batchMode_ = TritonBatchMode::Ragged; - outerDim_ = 1; - } -} - -void TritonClient::reset() { - if (!manualBatchMode_) - batchMode_ = TritonBatchMode::Rectangular; - for (auto& element : input_) { - element.second.reset(); - } - for (auto& element : output_) { - element.second.reset(); - } -} - -template -bool TritonClient::handle_exception(F&& call) { - //caught exceptions will be propagated to edm::WaitingTaskWithArenaHolder - CMS_SA_ALLOW try { - call(); - return true; - } - //TritonExceptions are intended/expected to be recoverable, i.e. retries should be allowed - catch (TritonException& e) { - e.convertToWarning(); - finish(false); - return false; - } - //other exceptions are not: execution should stop if they are encountered - catch (...) { - finish(false, std::current_exception()); - return false; - } -} - -const TritonService* TritonClient::service() const { - edm::ServiceRegistry::Operate op(token_); - edm::Service ts; - return &(*ts); -} - -const TritonService* TritonClient::localService() const { return isLocal_ ? service() : nullptr; } - -void TritonClient::getResults(const std::vector>& results) { - for (unsigned i = 0; i < results.size(); ++i) { - const auto& result = results[i]; - for (auto& [oname, output] : output_) { - //set shape here before output becomes const - if (output.variableDims()) { - std::vector tmp_shape; - TRITON_THROW_IF_ERROR(result->Shape(oname, &tmp_shape), - "getResults(): unable to get output shape for " + oname); - if (!noOuterDim_) - tmp_shape.erase(tmp_shape.begin()); - output.setShape(tmp_shape, i); - } - //extend lifetime - output.setResult(result, i); - //compute size after getting all result entries - if (i == results.size() - 1) - output.computeSizes(); - } - } -} - -//default case for sync and pseudo async -void TritonClient::evaluate() { - //undo previous signal from TritonException - if (tries_ > 0) { - // If we are retrying then the evaluate method is called outside the frameworks TBB thread pool. - // So we need to setup the service token for the current thread to access the service registry. - edm::ServiceRegistry::Operate op(token_); - edm::Service ts; - ts->notifyCallStatus(true); - } - - //in case there is nothing to process - if (batchSize() == 0) { - //call getResults on an empty vector - std::vector> empty_results; - getResults(empty_results); - finish(true); - return; - } - - //set up input pointers for triton (generalized for multi-request ragged batching case) - //one vector per request - unsigned nEntriesVal = nEntries(); - std::vector> inputsTriton(nEntriesVal); - for (auto& inputTriton : inputsTriton) { - inputTriton.reserve(input_.size()); - } - for (auto& [iname, input] : input_) { - for (unsigned i = 0; i < nEntriesVal; ++i) { - inputsTriton[i].push_back(input.data(i)); - } - } - - //set up output pointers similarly - std::vector> outputsTriton(nEntriesVal); - for (auto& outputTriton : outputsTriton) { - outputTriton.reserve(output_.size()); - } - for (auto& [oname, output] : output_) { - for (unsigned i = 0; i < nEntriesVal; ++i) { - outputsTriton[i].push_back(output.data(i)); - } - } - - //set up shared memory for output - auto success = handle_exception([&]() { - for (auto& element : output_) { - element.second.prepare(); - } - }); - if (!success) - return; - - // Get the status of the server prior to the request being made. - inference::ModelStatistics start_status; - success = handle_exception([&]() { - if (verbose()) - start_status = getServerSideStatus(); - }); - if (!success) - return; - - if (mode_ == SonicMode::Async) { - //non-blocking call - success = handle_exception([&]() { - TRITON_THROW_IF_ERROR(client_->AsyncInferMulti( - [start_status, this](std::vector resultsTmp) { - //immediately convert to shared_ptr - const auto& results = convertToShared(resultsTmp); - //check results - for (auto ptr : results) { - auto success = handle_exception([&]() { - TRITON_THROW_IF_ERROR( - ptr->RequestStatus(), "evaluate(): unable to get result(s)", localService()); - }); - if (!success) - return; - } - - if (verbose()) { - inference::ModelStatistics end_status; - auto success = handle_exception([&]() { end_status = getServerSideStatus(); }); - if (!success) - return; - - const auto& stats = summarizeServerStats(start_status, end_status); - reportServerSideStats(stats); - } - - //check result - auto success = handle_exception([&]() { getResults(results); }); - if (!success) - return; - - //finish - finish(true); - }, - options_, - inputsTriton, - outputsTriton, - headers_, - compressionAlgo_), - "evaluate(): unable to launch async run", - localService()); - }); - if (!success) - return; - } else { - //blocking call - std::vector resultsTmp; - success = handle_exception([&]() { - TRITON_THROW_IF_ERROR( - client_->InferMulti(&resultsTmp, options_, inputsTriton, outputsTriton, headers_, compressionAlgo_), - "evaluate(): unable to run and/or get result", - localService()); - }); - //immediately convert to shared_ptr - const auto& results = convertToShared(resultsTmp); - if (!success) - return; - - if (verbose()) { - inference::ModelStatistics end_status; - success = handle_exception([&]() { end_status = getServerSideStatus(); }); - if (!success) - return; - - const auto& stats = summarizeServerStats(start_status, end_status); - reportServerSideStats(stats); - } - - success = handle_exception([&]() { getResults(results); }); - if (!success) - return; - - finish(true); - } -} - -void TritonClient::reportServerSideStats(const TritonClient::ServerSideStats& stats) const { - std::stringstream msg; - - // https://github.com/triton-inference-server/server/blob/v2.3.0/src/clients/c++/perf_client/inference_profiler.cc - const uint64_t count = stats.success_count_; - msg << " Inference count: " << stats.inference_count_ << "\n"; - msg << " Execution count: " << stats.execution_count_ << "\n"; - msg << " Successful request count: " << count << "\n"; - - if (count > 0) { - auto get_avg_us = [count](uint64_t tval) { - constexpr uint64_t us_to_ns = 1000; - return tval / us_to_ns / count; - }; - - const uint64_t cumm_avg_us = get_avg_us(stats.cumm_time_ns_); - const uint64_t queue_avg_us = get_avg_us(stats.queue_time_ns_); - const uint64_t compute_input_avg_us = get_avg_us(stats.compute_input_time_ns_); - const uint64_t compute_infer_avg_us = get_avg_us(stats.compute_infer_time_ns_); - const uint64_t compute_output_avg_us = get_avg_us(stats.compute_output_time_ns_); - const uint64_t compute_avg_us = compute_input_avg_us + compute_infer_avg_us + compute_output_avg_us; - const uint64_t overhead = - (cumm_avg_us > queue_avg_us + compute_avg_us) ? (cumm_avg_us - queue_avg_us - compute_avg_us) : 0; - - msg << " Avg request latency: " << cumm_avg_us << " usec" - << "\n" - << " (overhead " << overhead << " usec + " - << "queue " << queue_avg_us << " usec + " - << "compute input " << compute_input_avg_us << " usec + " - << "compute infer " << compute_infer_avg_us << " usec + " - << "compute output " << compute_output_avg_us << " usec)" << std::endl; - } - - if (!debugName_.empty()) - edm::LogInfo(fullDebugName_) << msg.str(); -} - -TritonClient::ServerSideStats TritonClient::summarizeServerStats(const inference::ModelStatistics& start_status, - const inference::ModelStatistics& end_status) const { - TritonClient::ServerSideStats server_stats; - - server_stats.inference_count_ = end_status.inference_count() - start_status.inference_count(); - server_stats.execution_count_ = end_status.execution_count() - start_status.execution_count(); - server_stats.success_count_ = - end_status.inference_stats().success().count() - start_status.inference_stats().success().count(); - server_stats.cumm_time_ns_ = - end_status.inference_stats().success().ns() - start_status.inference_stats().success().ns(); - server_stats.queue_time_ns_ = end_status.inference_stats().queue().ns() - start_status.inference_stats().queue().ns(); - server_stats.compute_input_time_ns_ = - end_status.inference_stats().compute_input().ns() - start_status.inference_stats().compute_input().ns(); - server_stats.compute_infer_time_ns_ = - end_status.inference_stats().compute_infer().ns() - start_status.inference_stats().compute_infer().ns(); - server_stats.compute_output_time_ns_ = - end_status.inference_stats().compute_output().ns() - start_status.inference_stats().compute_output().ns(); - - return server_stats; -} - -inference::ModelStatistics TritonClient::getServerSideStatus() const { - if (verbose_) { - inference::ModelStatisticsResponse resp; - TRITON_THROW_IF_ERROR(client_->ModelInferenceStatistics(&resp, options_[0].model_name_, options_[0].model_version_), - "getServerSideStatus(): unable to get model statistics", - localService()); - return *(resp.model_stats().begin()); - } - return inference::ModelStatistics{}; -} - -//for fillDescriptions -void TritonClient::fillPSetDescription(edm::ParameterSetDescription& iDesc) { - edm::ParameterSetDescription descClient; - fillBasePSetDescription(descClient); - descClient.add("modelName"); - descClient.add("modelVersion", ""); - descClient.add("modelConfigPath"); - //server parameters should not affect the physics results - descClient.addUntracked("preferredServer", ""); - descClient.addUntracked("timeout"); - descClient.ifValue(edm::ParameterDescription("timeoutUnit", "seconds", false), - edm::allowedValues("seconds", "milliseconds", "microseconds")); - descClient.addUntracked("useSharedMemory", true); - descClient.addUntracked("compression", ""); - descClient.addUntracked>("outputs", {}); - iDesc.add("Client", descClient); -} +#include "FWCore/Concurrency/interface/WaitingTask.h" +#include "FWCore/MessageLogger/interface/MessageLogger.h" +#include "FWCore/ParameterSet/interface/FileInPath.h" +#include "FWCore/ParameterSet/interface/allowedValues.h" +#include "FWCore/ServiceRegistry/interface/Service.h" +#include "FWCore/Utilities/interface/Exception.h" +#include "HeterogeneousCore/SonicTriton/interface/TritonClient.h" +#include "HeterogeneousCore/SonicTriton/interface/TritonException.h" +#include "HeterogeneousCore/SonicTriton/interface/TritonService.h" +#include "HeterogeneousCore/SonicTriton/interface/triton_utils.h" + +#include "grpc_client.h" +#include "grpc_service.pb.h" +#include "model_config.pb.h" + +#include "google/protobuf/text_format.h" +#include "google/protobuf/io/zero_copy_stream_impl.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace tc = triton::client; + +namespace { + // Minimal ParameterSet to satisfy SonicClientBase requirements during unit tests + edm::ParameterSet makeMinimalSonicParamsForTest() { + edm::ParameterSet params; + params.addParameter("mode", "PseudoAsync"); + + edm::ParameterSet defaultRetry; + defaultRetry.addParameter("retryType", "RetrySameServerAction"); + defaultRetry.addUntrackedParameter("allowedTries", 0u); + std::vector retryVec{defaultRetry}; + params.addParameter>("Retry", retryVec); + + return params; + } + grpc_compression_algorithm getCompressionAlgo(const std::string& name) { + if (name.empty() or name.compare("none") == 0) + return grpc_compression_algorithm::GRPC_COMPRESS_NONE; + else if (name.compare("deflate") == 0) + return grpc_compression_algorithm::GRPC_COMPRESS_DEFLATE; + else if (name.compare("gzip") == 0) + return grpc_compression_algorithm::GRPC_COMPRESS_GZIP; + else + throw cms::Exception("GrpcCompression") + << "Unknown compression algorithm requested: " << name << " (choices: none, deflate, gzip)"; + } + + std::vector> convertToShared(const std::vector& tmp) { + std::vector> results; + results.reserve(tmp.size()); + std::transform(tmp.begin(), tmp.end(), std::back_inserter(results), [](tc::InferResult* ptr) { + return std::shared_ptr(ptr); + }); + return results; + } +} // namespace + +//based on https://github.com/triton-inference-server/server/blob/v2.3.0/src/clients/c++/examples/simple_grpc_async_infer_client.cc +//and https://github.com/triton-inference-server/server/blob/v2.3.0/src/clients/c++/perf_client/perf_client.cc + +TritonClient::TritonClient(const edm::ParameterSet& params, const std::string& debugName) + : SonicClient(params, debugName, "TritonClient"), + batchMode_(TritonBatchMode::Rectangular), + manualBatchMode_(false), + verbose_(params.getUntrackedParameter("verbose")), + useSharedMemory_(params.getUntrackedParameter("useSharedMemory")), + compressionAlgo_(getCompressionAlgo(params.getUntrackedParameter("compression"))) { + options_.emplace_back(params.getParameter("modelName")); + + edm::Service ts; + + // We save the token to be able to notify the service in case of an exception in the evaluate method. + // The evaluate method can be called outside the frameworks TBB threadpool in the case of a retry. In + // this case the context is not setup to access the service registry, we need the service token to + // create the context. + token_ = edm::ServiceRegistry::instance().presentToken(); + + //Connect to server + updateServer(params.getUntrackedParameter("preferredServer")); + + //set options + options_[0].model_version_ = params.getParameter("modelVersion"); + options_[0].client_timeout_ = params.getUntrackedParameter("timeout"); + //convert to microseconds + const auto& timeoutUnit = params.getUntrackedParameter("timeoutUnit"); + unsigned conversion = 1; + if (timeoutUnit == "seconds") + conversion = 1e6; + else if (timeoutUnit == "milliseconds") + conversion = 1e3; + else if (timeoutUnit == "microseconds") + conversion = 1; + else + throw cms::Exception("Configuration") << "Unknown timeout unit: " << timeoutUnit; + options_[0].client_timeout_ *= conversion; + + //get fixed parameters from local config + inference::ModelConfig localModelConfig; + { + const std::string localModelConfigPath(params.getParameter("modelConfigPath").fullPath()); + int fileDescriptor = open(localModelConfigPath.c_str(), O_RDONLY); + if (fileDescriptor < 0) + throw TritonException("LocalFailure") + << "TritonClient(): unable to open local model config: " << localModelConfigPath; + google::protobuf::io::FileInputStream localModelConfigInput(fileDescriptor); + localModelConfigInput.SetCloseOnDelete(true); + if (!google::protobuf::TextFormat::Parse(&localModelConfigInput, &localModelConfig)) + throw TritonException("LocalFailure") + << "TritonClient(): unable to parse local model config: " << localModelConfigPath; + } + + //check batch size limitations (after i/o setup) + //triton uses max batch size = 0 to denote a model that does not support native batching (using the outer dimension) + //but for models that do support batching (native or otherwise), a given event may set batch size 0 to indicate no valid input is present + //so set the local max to 1 and keep track of "no outer dim" case + maxOuterDim_ = localModelConfig.max_batch_size(); + noOuterDim_ = maxOuterDim_ == 0; + maxOuterDim_ = std::max(1u, maxOuterDim_); + //propagate batch size + setBatchSize(1); + + //compare model checksums to remote config to enforce versioning + inference::ModelConfigResponse modelConfigResponse; + TRITON_THROW_IF_ERROR(client_->ModelConfig(&modelConfigResponse, options_[0].model_name_, options_[0].model_version_), + "TritonClient(): unable to get model config", + localService()); + inference::ModelConfig remoteModelConfig(modelConfigResponse.config()); + + std::map> checksums; + size_t fileCounter = 0; + for (const auto& modelConfig : {localModelConfig, remoteModelConfig}) { + const auto& agents = modelConfig.model_repository_agents().agents(); + auto agent = std::find_if(agents.begin(), agents.end(), [](auto const& a) { return a.name() == "checksum"; }); + if (agent != agents.end()) { + const auto& params = agent->parameters(); + for (const auto& [key, val] : params) { + // only check the requested version + if (key.compare(0, options_[0].model_version_.size() + 1, options_[0].model_version_ + "/") == 0) + checksums[key][fileCounter] = val; + } + } + ++fileCounter; + } + std::vector incorrect; + for (const auto& [key, val] : checksums) { + if (checksums[key][0] != checksums[key][1]) + incorrect.push_back(key); + } + if (!incorrect.empty()) + throw TritonException("ModelVersioning") << "The following files have incorrect checksums on the remote server: " + << triton_utils::printColl(incorrect, ", "); + + //get model info + inference::ModelMetadataResponse modelMetadata; + TRITON_THROW_IF_ERROR(client_->ModelMetadata(&modelMetadata, options_[0].model_name_, options_[0].model_version_), + "TritonClient(): unable to get model metadata", + localService()); + + //get input and output (which know their sizes) + const auto& nicInputs = modelMetadata.inputs(); + const auto& nicOutputs = modelMetadata.outputs(); + + //report all model errors at once + std::stringstream msg; + std::string msg_str; + + //currently no use case is foreseen for a model with zero inputs or outputs + if (nicInputs.empty()) + msg << "Model on server appears malformed (zero inputs)\n"; + + if (nicOutputs.empty()) + msg << "Model on server appears malformed (zero outputs)\n"; + + //stop if errors + msg_str = msg.str(); + if (!msg_str.empty()) + throw cms::Exception("ModelErrors") << msg_str; + + //setup input map + std::stringstream io_msg; + if (verbose_) + io_msg << "Model inputs: " + << "\n"; + for (const auto& nicInput : nicInputs) { + const auto& iname = nicInput.name(); + auto [curr_itr, success] = input_.emplace(std::piecewise_construct, + std::forward_as_tuple(iname), + std::forward_as_tuple(iname, nicInput, this, ts->pid())); + auto& curr_input = curr_itr->second; + if (verbose_) { + io_msg << " " << iname << " (" << curr_input.dname() << ", " << curr_input.byteSize() + << " b) : " << triton_utils::printColl(curr_input.shape()) << "\n"; + } + } + + //allow selecting only some outputs from server + const auto& v_outputs = params.getUntrackedParameter>("outputs"); + std::unordered_set s_outputs(v_outputs.begin(), v_outputs.end()); + + //setup output map + if (verbose_) + io_msg << "Model outputs: " + << "\n"; + for (const auto& nicOutput : nicOutputs) { + const auto& oname = nicOutput.name(); + if (!s_outputs.empty() and s_outputs.find(oname) == s_outputs.end()) + continue; + auto [curr_itr, success] = output_.emplace(std::piecewise_construct, + std::forward_as_tuple(oname), + std::forward_as_tuple(oname, nicOutput, this, ts->pid())); + auto& curr_output = curr_itr->second; + if (verbose_) { + io_msg << " " << oname << " (" << curr_output.dname() << ", " << curr_output.byteSize() + << " b) : " << triton_utils::printColl(curr_output.shape()) << "\n"; + } + if (!s_outputs.empty()) + s_outputs.erase(oname); + } + + //check if any requested outputs were not available + if (!s_outputs.empty()) + throw cms::Exception("MissingOutput") + << "Some requested outputs were not available on the server: " << triton_utils::printColl(s_outputs); + + //print model info + std::stringstream model_msg; + if (verbose_) { + model_msg << "Model name: " << options_[0].model_name_ << "\n" + << "Model version: " << options_[0].model_version_ << "\n" + << "Model max outer dim: " << (noOuterDim_ ? 0 : maxOuterDim_) << "\n"; + edm::LogInfo(fullDebugName_) << model_msg.str() << io_msg.str(); + } +} + +TritonClient::~TritonClient() { + //by default: members of this class destroyed before members of base class + //in shared memory case, TritonMemResource (member of TritonData) unregisters from client_ in its destructor + //but input/output objects are member of base class, so destroyed after client_ (member of this class) + //therefore, clear the maps here + input_.clear(); + output_.clear(); +} + +void TritonClient::setBatchMode(TritonBatchMode batchMode) { + unsigned oldBatchSize = batchSize(); + batchMode_ = batchMode; + manualBatchMode_ = true; + //this allows calling setBatchSize() and setBatchMode() in either order consistently to change back and forth + //includes handling of change from ragged to rectangular if multiple entries already created + setBatchSize(oldBatchSize); +} + +void TritonClient::resetBatchMode() { + batchMode_ = TritonBatchMode::Rectangular; + manualBatchMode_ = false; +} + +unsigned TritonClient::nEntries() const { return !input_.empty() ? input_.begin()->second.entries_.size() : 0; } + +unsigned TritonClient::batchSize() const { return batchMode_ == TritonBatchMode::Rectangular ? outerDim_ : nEntries(); } + +bool TritonClient::setBatchSize(unsigned bsize) { + if (batchMode_ == TritonBatchMode::Rectangular) { + if (bsize > maxOuterDim_) { + throw TritonException("LocalFailure") + << "Requested batch size " << bsize << " exceeds server-specified max batch size " << maxOuterDim_ << "."; + return false; + } else { + outerDim_ = bsize; + //take min to allow resizing to 0 + resizeEntries(std::min(outerDim_, 1u)); + return true; + } + } else { + resizeEntries(bsize); + outerDim_ = 1; + return true; + } +} + +void TritonClient::resizeEntries(unsigned entry) { + if (entry > nEntries()) + //addEntry(entry) extends the vector to size entry+1 + addEntry(entry - 1); + else if (entry < nEntries()) { + for (auto& element : input_) { + element.second.entries_.resize(entry); + } + for (auto& element : output_) { + element.second.entries_.resize(entry); + } + } +} + +void TritonClient::addEntry(unsigned entry) { + for (auto& element : input_) { + element.second.addEntryImpl(entry); + } + for (auto& element : output_) { + element.second.addEntryImpl(entry); + } + if (entry > 0) { + batchMode_ = TritonBatchMode::Ragged; + outerDim_ = 1; + } +} + +void TritonClient::reset() { + if (!manualBatchMode_) + batchMode_ = TritonBatchMode::Rectangular; + for (auto& element : input_) { + element.second.reset(); + } + for (auto& element : output_) { + element.second.reset(); + } +} + +template +bool TritonClient::handle_exception(F&& call) { + //caught exceptions will be propagated to edm::WaitingTaskWithArenaHolder + CMS_SA_ALLOW try { + call(); + return true; + } + //TritonExceptions are intended/expected to be recoverable, i.e. retries should be allowed + catch (TritonException& e) { + e.convertToWarning(); + finish(false); + return false; + } + //other exceptions are not: execution should stop if they are encountered + catch (...) { + finish(false, std::current_exception()); + return false; + } +} + +template +bool TritonClient::handle_exception_holder(edm::WaitingTaskWithArenaHolder& fh, F&& call) { + CMS_SA_ALLOW try { + call(); + return true; + } catch (TritonException& e) { + e.convertToWarning(); + inferSuccess_ = false; + fh.doneWaiting(nullptr); // retryable: schedules finish(false) on TBB + return false; + } catch (...) { + fh.doneWaiting(std::current_exception()); // non-retryable + return false; + } +} + +const TritonService* TritonClient::service() const { + edm::ServiceRegistry::Operate op(token_); + edm::Service ts; + return &(*ts); +} + +const TritonService* TritonClient::localService() const { return isLocal_ ? service() : nullptr; } + +void TritonClient::getResults(const std::vector>& results) { + for (unsigned i = 0; i < results.size(); ++i) { + const auto& result = results[i]; + for (auto& [oname, output] : output_) { + //set shape here before output becomes const + if (output.variableDims()) { + std::vector tmp_shape; + TRITON_THROW_IF_ERROR(result->Shape(oname, &tmp_shape), + "getResults(): unable to get output shape for " + oname); + if (!noOuterDim_) + tmp_shape.erase(tmp_shape.begin()); + output.setShape(tmp_shape, i); + } + //extend lifetime + output.setResult(result, i); + //compute size after getting all result entries + if (i == results.size() - 1) + output.computeSizes(); + } + } +} + +//default case for sync and pseudo async +void TritonClient::evaluate() { + //undo previous signal from TritonException + if (totalTries_ > 0) { + // If we are retrying then the evaluate method is called outside the frameworks TBB thread pool. + // So we need to setup the service token for the current thread to access the service registry. + edm::ServiceRegistry::Operate op(token_); + edm::Service ts; + ts->notifyCallStatus(true); + } + + //in case there is nothing to process + if (batchSize() == 0) { + //call getResults on an empty vector + std::vector> empty_results; + getResults(empty_results); + finish(true); + return; + } + + //set up input pointers for triton (generalized for multi-request ragged batching case) + //one vector per request + unsigned nEntriesVal = nEntries(); + std::vector> inputsTriton(nEntriesVal); + for (auto& inputTriton : inputsTriton) { + inputTriton.reserve(input_.size()); + } + for (auto& [iname, input] : input_) { + for (unsigned i = 0; i < nEntriesVal; ++i) { + inputsTriton[i].push_back(input.data(i)); + } + } + + //set up output pointers similarly + std::vector> outputsTriton(nEntriesVal); + for (auto& outputTriton : outputsTriton) { + outputTriton.reserve(output_.size()); + } + for (auto& [oname, output] : output_) { + for (unsigned i = 0; i < nEntriesVal; ++i) { + outputsTriton[i].push_back(output.data(i)); + } + } + + //set up shared memory for output + auto success = handle_exception([&]() { + for (auto& element : output_) { + element.second.prepare(); + } + }); + //edm::LogInfo("TritonClient") << "evaluate() return 1"; + if (!success) + return; + + // Get the status of the server prior to the request being made. + inference::ModelStatistics start_status; + success = handle_exception([&]() { + if (verbose()) + start_status = getServerSideStatus(); + }); + //edm::LogInfo("TritonClient") << "evaluate() return 2"; + if (!success) + return; + + if (mode_ == SonicMode::Async) { + // Reset before each inference attempt: false=retryable/not-yet-succeeded. + // The callback sets this to true on success before calling doneWaiting(nullptr). + // If the launch throws (lambda destroyed inside AsyncInferMulti), the holder destructor + // calls doneWaiting(nullptr) with inferSuccess_=false → finish(false) [retryable] on TBB. + inferSuccess_ = false; + + // Create holder[finish]: when doneWaiting() is called, finish() is scheduled as a TBB task + // so retry logic (finish→retry→updateServer) never runs on the gRPC callback thread. + auto* finishTask = edm::make_waiting_task([this](std::exception_ptr const* excptr) { + if (excptr) + finish(false, *excptr); // non-retryable exception + else + finish(inferSuccess_.load()); // true=success, false=retryable failure + }); + edm::WaitingTaskWithArenaHolder finishHolder(*holder_->group(), finishTask); + + // Launch async inference. + // On TritonException: fh destructor (in destroyed lambda) calls doneWaiting(nullptr) + // → schedules finish(false) [retryable] on TBB. Do NOT call finish() directly here. + CMS_SA_ALLOW try { + TRITON_THROW_IF_ERROR( + client_->AsyncInferMulti( + [start_status, fh = std::move(finishHolder), this](std::vector resultsTmp) mutable { + //immediately convert to shared_ptr + const auto& results = convertToShared(resultsTmp); + //check results + for (auto ptr : results) { + auto ok = handle_exception_holder(fh, [&]() { + TRITON_THROW_IF_ERROR(ptr->RequestStatus(), "evaluate(): unable to get result(s)", localService()); + }); + if (!ok) + return; + } + + if (verbose()) { + inference::ModelStatistics end_status; + auto ok = handle_exception_holder(fh, [&]() { end_status = getServerSideStatus(); }); + if (!ok) + return; + const auto& stats = summarizeServerStats(start_status, end_status); + reportServerSideStats(stats); + } + + auto ok = handle_exception_holder(fh, [&]() { getResults(results); }); + if (!ok) + return; + + inferSuccess_ = true; + fh.doneWaiting(nullptr); // success: schedules finish(true) on TBB + }, + options_, + inputsTriton, + outputsTriton, + headers_, + compressionAlgo_), + "evaluate(): unable to launch async run", + localService()); + } catch (TritonException& e) { + e.convertToWarning(); + // fh destructor already called doneWaiting(nullptr) → finish(false) scheduled on TBB + } + } else { + //edm::LogInfo("TritonClient") << "evaluate() return 4"; + //blocking call + std::vector resultsTmp; + success = handle_exception([&]() { + TRITON_THROW_IF_ERROR( + client_->InferMulti(&resultsTmp, options_, inputsTriton, outputsTriton, headers_, compressionAlgo_), + "evaluate(): unable to run and/or get result", + localService()); + }); + //immediately convert to shared_ptr + const auto& results = convertToShared(resultsTmp); + if (!success) + return; + + if (verbose()) { + inference::ModelStatistics end_status; + success = handle_exception([&]() { end_status = getServerSideStatus(); }); + if (!success) + return; + + const auto& stats = summarizeServerStats(start_status, end_status); + reportServerSideStats(stats); + } + + success = handle_exception([&]() { getResults(results); }); + if (!success) + return; + + finish(true); + } +} + +void TritonClient::reportServerSideStats(const TritonClient::ServerSideStats& stats) const { + std::stringstream msg; + + // https://github.com/triton-inference-server/server/blob/v2.3.0/src/clients/c++/perf_client/inference_profiler.cc + const uint64_t count = stats.success_count_; + msg << " Inference count: " << stats.inference_count_ << "\n"; + msg << " Execution count: " << stats.execution_count_ << "\n"; + msg << " Successful request count: " << count << "\n"; + + if (count > 0) { + auto get_avg_us = [count](uint64_t tval) { + constexpr uint64_t us_to_ns = 1000; + return tval / us_to_ns / count; + }; + + const uint64_t cumm_avg_us = get_avg_us(stats.cumm_time_ns_); + const uint64_t queue_avg_us = get_avg_us(stats.queue_time_ns_); + const uint64_t compute_input_avg_us = get_avg_us(stats.compute_input_time_ns_); + const uint64_t compute_infer_avg_us = get_avg_us(stats.compute_infer_time_ns_); + const uint64_t compute_output_avg_us = get_avg_us(stats.compute_output_time_ns_); + const uint64_t compute_avg_us = compute_input_avg_us + compute_infer_avg_us + compute_output_avg_us; + const uint64_t overhead = + (cumm_avg_us > queue_avg_us + compute_avg_us) ? (cumm_avg_us - queue_avg_us - compute_avg_us) : 0; + + msg << " Avg request latency: " << cumm_avg_us << " usec" + << "\n" + << " (overhead " << overhead << " usec + " + << "queue " << queue_avg_us << " usec + " + << "compute input " << compute_input_avg_us << " usec + " + << "compute infer " << compute_infer_avg_us << " usec + " + << "compute output " << compute_output_avg_us << " usec)" << std::endl; + } + + if (!debugName_.empty()) + edm::LogInfo(fullDebugName_) << msg.str(); +} + +TritonClient::ServerSideStats TritonClient::summarizeServerStats(const inference::ModelStatistics& start_status, + const inference::ModelStatistics& end_status) const { + TritonClient::ServerSideStats server_stats; + + server_stats.inference_count_ = end_status.inference_count() - start_status.inference_count(); + server_stats.execution_count_ = end_status.execution_count() - start_status.execution_count(); + server_stats.success_count_ = + end_status.inference_stats().success().count() - start_status.inference_stats().success().count(); + server_stats.cumm_time_ns_ = + end_status.inference_stats().success().ns() - start_status.inference_stats().success().ns(); + server_stats.queue_time_ns_ = end_status.inference_stats().queue().ns() - start_status.inference_stats().queue().ns(); + server_stats.compute_input_time_ns_ = + end_status.inference_stats().compute_input().ns() - start_status.inference_stats().compute_input().ns(); + server_stats.compute_infer_time_ns_ = + end_status.inference_stats().compute_infer().ns() - start_status.inference_stats().compute_infer().ns(); + server_stats.compute_output_time_ns_ = + end_status.inference_stats().compute_output().ns() - start_status.inference_stats().compute_output().ns(); + + return server_stats; +} + +inference::ModelStatistics TritonClient::getServerSideStatus() const { + if (verbose_) { + inference::ModelStatisticsResponse resp; + TRITON_THROW_IF_ERROR(client_->ModelInferenceStatistics(&resp, options_[0].model_name_, options_[0].model_version_), + "getServerSideStatus(): unable to get model statistics", + localService()); + return *(resp.model_stats().begin()); + } + return inference::ModelStatistics{}; +} + +void TritonClient::updateServer(const std::string& serverName) { + //get appropriate server for this model + auto ts = service(); + + const auto& serverMap = ts->resolveServer(options_[0].model_name_, serverName); + + const auto& server = serverMap.second; + + //update server name + serverName_ = serverMap.first; + + serverType_ = server.type; + edm::LogInfo("TritonDiscovery") << debugName_ << " assigned server: " << server.url; + //enforce sync mode for fallback CPU server to avoid contention + //todo: could enforce async mode otherwise (unless mode was specified by user?) + if (serverType_ == TritonServerType::LocalCPU) + setMode(SonicMode::Sync); + isLocal_ = serverType_ == TritonServerType::LocalCPU or serverType_ == TritonServerType::LocalGPU; + + // updateServer() is always called from a TBB thread (via finish() -> retry() -> updateServer()), + // so it is safe to destroy the old gRPC client immediately. + client_.reset(); + + //connect to the server + TRITON_THROW_IF_ERROR( + tc::InferenceServerGrpcClient::Create(&client_, server.url, false, server.useSsl, server.sslOptions), + "TritonClient(): unable to create inference context", + localService()); +} + +void TritonClient::switchToFallback() { + edm::ServiceRegistry::Operate op(token_); + edm::Service ts; + + // Start the fallback server if it has not been started yet (idempotent). + ts->startFallbackServer(); + if (!ts->fallbackStarted()) + throw TritonException("LocalFailure") + << "TritonClient::switchToFallback: fallback server is not available " + "(check that TritonService fallback.enable = True and a model path is configured)"; + + // Dynamically load this client's model onto the fallback server. + ts->loadModel(options_[0].model_name_); + + // Re-point the gRPC connection at the fallback server. + updateServer(TritonService::Server::fallbackName); +} + +//for fillDescriptions +void TritonClient::fillPSetDescription(edm::ParameterSetDescription& iDesc) { + edm::ParameterSetDescription descClient; + fillBasePSetDescription(descClient); + descClient.add("modelName"); + descClient.add("modelVersion", ""); + descClient.add("modelConfigPath"); + //server parameters should not affect the physics results + descClient.addUntracked("preferredServer", ""); + descClient.addUntracked("timeout"); + descClient.ifValue(edm::ParameterDescription("timeoutUnit", "seconds", false), + edm::allowedValues("seconds", "milliseconds", "microseconds")); + descClient.addUntracked("useSharedMemory", true); + descClient.addUntracked("compression", ""); + descClient.addUntracked>("outputs", {}); + iDesc.add("Client", descClient); +} + +void TritonClient::connectToServer(const std::string& url) { + // Update client state for a generic remote server + serverType_ = TritonServerType::Remote; + isLocal_ = false; + + edm::LogInfo("TritonDiscovery") << debugName_ << " connecting to server: " << url; + + // Use default SSL options + triton::client::SslOptions sslOptions; + bool useSsl = false; // Assuming no SSL for direct URL connection + + // Connect to the server + TRITON_THROW_IF_ERROR(triton::client::InferenceServerGrpcClient::Create(&client_, url, false, useSsl, sslOptions), + "TritonClient::connectToServer(): unable to create inference context"); +} + +//constructor for testing +TritonClient::TritonClient() : SonicClient(makeMinimalSonicParamsForTest(), "TritonClient_test", "TritonClient") {} diff --git a/HeterogeneousCore/SonicTriton/src/TritonService.cc b/HeterogeneousCore/SonicTriton/src/TritonService.cc index 112e70d1d1ea2..b10ffacadcaa5 100644 --- a/HeterogeneousCore/SonicTriton/src/TritonService.cc +++ b/HeterogeneousCore/SonicTriton/src/TritonService.cc @@ -120,11 +120,14 @@ TritonService::TritonService(const edm::ParameterSet& pset, edm::ActivityRegistr << "TritonService: Not allowed to specify more than one server with same name (" << serverName << ")"; } - //loop over all servers: check which models they have + //loop over all servers: check which models they have, populate serverHealth std::string msg; if (verbose_) msg = "List of models for each server:\n"; for (auto& [serverName, server] : servers_) { + //populate serverHealth + serversHealth_.emplace(serverName, ServerHealth{}); + std::unique_ptr client; TRITON_THROW_IF_ERROR( tc::InferenceServerGrpcClient::Create(&client, server.url, false, server.useSsl, server.sslOptions), @@ -138,7 +141,8 @@ TritonService::TritonService(const edm::ParameterSet& pset, edm::ActivityRegistr edm::LogInfo("TritonService") << "Server " << serverName << ": url = " << server.url << ", version = " << serverMetaResponse.version(); else - edm::LogInfo("TritonService") << "unable to get metadata for " + serverName + " (" + server.url + ")"; + edm::LogInfo("TritonService") << "unable to get metadata for " + serverName + " (" + server.url + ")" + << err.Message(); } //if this query fails, it indicates that the server is nonresponsive or saturated @@ -191,64 +195,214 @@ void TritonService::addModel(const std::string& modelName, const std::string& pa if (!allowAddModel_) throw cms::Exception("DisallowedAddModel") << "TritonService: Attempt to call addModel() outside of module constructors"; - //if model is not in the list, then no specified server provides it - auto mit = models_.find(modelName); - if (mit == models_.end()) { - auto& modelInfo(unservedModels_.emplace(modelName, path).first->second); - modelInfo.modules.insert(currentModuleId_); - //only keep track of modules that need unserved models - modules_.emplace(currentModuleId_, modelName); - } + + auto& modelInfo(models_.emplace(modelName, path).first->second); + // Update path if model was previously added (e.g., by server scanning) with empty path + if (modelInfo.path.empty() && !path.empty()) + modelInfo.path = path; + + modelInfo.modules.insert(currentModuleId_); + modules_.emplace(currentModuleId_, modelName); } void TritonService::postModuleConstruction(edm::ModuleDescription const& desc) { allowAddModel_ = false; } void TritonService::preModuleDestruction(edm::ModuleDescription const& desc) { - //remove destructed modules from unserved list - if (unservedModels_.empty()) - return; auto id = desc.id(); auto oit = modules_.find(id); if (oit != modules_.end()) { const auto& moduleInfo(oit->second); - auto mit = unservedModels_.find(moduleInfo.model); - if (mit != unservedModels_.end()) { + auto mit = models_.find(moduleInfo.model); + if (mit != models_.end()) { auto& modelInfo(mit->second); modelInfo.modules.erase(id); - //remove a model if it is no longer needed by any modules - if (modelInfo.modules.empty()) - unservedModels_.erase(mit); } modules_.erase(oit); } } -//second return value is only true if fallback CPU server is being used -TritonService::Server TritonService::serverInfo(const std::string& model, const std::string& preferred) const { +// Returns the name of the server assigned to serve the given model, or nullptr if no server is currently assigned. +// If a preferred server(current server) is specified but unavailable, falls back to any assigned server. +// Callers are responsible for handling the nullptr case. +const std::string* TritonService::resolveServerName(const std::string& model, const std::string& preferred) const { auto mit = models_.find(model); - if (mit == models_.end()) - throw cms::Exception("MissingModel") << "TritonService: There are no servers that provide model " << model; - const auto& modelInfo(mit->second); - const auto& modelServers = modelInfo.servers; + if (mit == models_.end() || mit->second.servers.empty()) + return nullptr; // no server assigned - caller decides what to do + + const auto& modelServers = mit->second.servers; - auto msit = modelServers.end(); if (!preferred.empty()) { - msit = modelServers.find(preferred); - //todo: add a "strict" parameter to stop execution if preferred server isn't found? - if (msit == modelServers.end()) - edm::LogWarning("PreferredServer") << "Preferred server " << preferred << " for model " << model - << " not available, will choose another server"; + auto msit = modelServers.find(preferred); + if (msit != modelServers.end()) + return &(*msit); + edm::LogWarning("PreferredServer") << "Preferred server " << preferred << " for model " << model + << " not available, will choose another server"; } - const auto& serverName(msit == modelServers.end() ? *modelServers.begin() : preferred); + //Prefer remote servers over fallback if available + if (modelServers.size() > 1) { + auto rit = std::find_if(modelServers.begin(), modelServers.end(), [this](const std::string& name) { + auto sit = servers_.find(name); + return sit != servers_.end() && !sit->second.isFallback; + }); + if (rit != modelServers.end()) + return &(*rit); + } + return &(*modelServers.begin()); +} + +// Returns the full server info for the server assigned to serve the given model. +// Throws MissingModel if no server is currently assigned. +// Wraps resolveServerName; use that directly if nullptr should be handled by the caller +const std::pair& TritonService::resolveServer( + const std::string& model, const std::string& preferred) const { + const auto* name = resolveServerName(model, preferred); + if (!name) + throw cms::Exception("MissingModel") << "TritonService: There are no servers that provide model " << model; + return *servers_.find(*name); +} - //todo: use some algorithm to select server rather than just picking arbitrarily - const auto& server(servers_.find(serverName)->second); - return server; +// Returns the list of model names that are not currently assigned to any server. +std::vector TritonService::unassignedModels() const { + std::vector result; + for (const auto& [name, info] : models_) { + if (info.servers.empty()) + result.push_back(name); + } + return result; } -void TritonService::preBeginJob(edm::ProcessContext const&) { - //only need fallback if there are unserved models - if (!fallbackOpts_.enable or unservedModels_.empty()) +void TritonService::updateServerHealth(const std::string& modelName) const { + for (auto& [serverName, server] : servers_) { + edm::LogInfo("TritonService") << "Updating server health for server = " << serverName; + if (server.isFallback) { + edm::LogInfo("TritonService") << serverName << " is skipped because it is a fallback server"; + continue; // fallback is a last resort, not a candidate for getBestServer + } + try { + std::unique_ptr client; + TRITON_THROW_IF_ERROR( + tc::InferenceServerGrpcClient::Create(&client, server.url, false, server.useSsl, server.sslOptions), + "TritonService(): unable to create inference context for " + serverName + " (" + server.url + ")"); + + bool live = false, ready = false; + TRITON_THROW_IF_ERROR(client->IsServerLive(&live), + "TritonService(): unable to query IsServerLive " + serverName + " (" + server.url + ")"); + TRITON_THROW_IF_ERROR(client->IsServerReady(&ready), + "TritonService(): unable to query IsServerReady " + serverName + " (" + server.url + ")"); + + edm::LogInfo("TritonService") << serverName << " : live = " << live << " ready = " << ready; + + inference::ModelStatisticsResponse stats; + if (!modelName.empty()) { + client->ModelInferenceStatistics(&stats, modelName); + } else { + for (const auto& m : server.models) { + client->ModelInferenceStatistics(&stats, m); + } + } + + uint64_t infer_count = 0, queue_count = 0, failures = 0; + double avgQueueTimeMs = 0.0; + double avgInferTimeMs = 0.0; + + for (const auto& mstat : stats.model_stats()) { + if (modelName.empty() || mstat.name() == modelName) { + const auto& infer = mstat.inference_stats(); + + infer_count += infer.compute_infer().count(); + avgInferTimeMs += infer.compute_infer().ns() / 1e3; + queue_count += infer.queue().count(); + avgQueueTimeMs += infer.queue().ns() / 1e3; + failures += infer.fail().count(); + } + } + // Update health map safely with accessor + tbb::concurrent_hash_map::accessor acc; + serversHealth_.find(acc, serverName); + + ServerHealth& health = acc->second; + health.live = live; + health.ready = ready; + health.failureCount = failures; + health.avgQueueTimeMs = (queue_count > 0) ? avgQueueTimeMs / queue_count : 0.0; + health.avgInferTimeMs = (infer_count > 0) ? avgInferTimeMs / infer_count : 0.0; + + } catch (const TritonException& e) { + // mark existing entry unhealthy if present + tbb::concurrent_hash_map::accessor acc; + if (serversHealth_.find(acc, serverName)) { + ServerHealth& health = acc->second; + health.live = false; + health.ready = false; + } + } catch (const std::exception& e) { + // fallback for other exceptions + tbb::concurrent_hash_map::accessor acc; + if (serversHealth_.find(acc, serverName)) { + ServerHealth& health = acc->second; + health.live = false; + health.ready = false; + } + } + } +} + +std::optional TritonService::getBestServer(const std::string& modelName, + const std::string& IgnoreServer) const { + std::optional bestServerName; + ServerHealth bestHealth; + + // get fresh ServerHealth statistics + updateServerHealth(modelName); + edm::LogInfo("TritonService") << "Getting best server"; + + for (auto& [serverName, server] : servers_) { + if (serverName == IgnoreServer) { + edm::LogInfo("TritonService") << serverName << " is ignored"; + continue; // skip ignored server + } + if (server.isFallback) { + edm::LogInfo("TritonService") << serverName << " is skipped because it is a fallback server"; + continue; // fallback is a last resort, not a candidate for getBestServer + } + if (server.models.find(modelName) == server.models.end()) { + edm::LogInfo("TritonService") << serverName << " is skipped because it does not have " << modelName; + continue; // server doesn't have model + } + + tbb::concurrent_hash_map::const_accessor acc; + if (!serversHealth_.find(acc, serverName)) { + edm::LogInfo("TritonService") << serverName << " is skipped because it does not have health info"; + continue; // no health info + } + + const ServerHealth& health = acc->second; + + if (!health.live || !health.ready) { + edm::LogInfo("TritonService") << serverName << " is skipped because is not live or ready"; + continue; // skip unhealthy + } + + // Select server according to rules: + // 1) lowest failureCount + // 2) tie-breaker: lowest avgQueueTimeMs + if (!bestServerName || health.failureCount < bestHealth.failureCount || + (health.failureCount == bestHealth.failureCount && health.avgQueueTimeMs < bestHealth.avgQueueTimeMs)) { + bestServerName = serverName; + bestHealth = health; + } + } + if (verbose_ && bestServerName) { + edm::LogInfo("TritonDiscovery") << "Chosen server for model '" << modelName << "': " << *bestServerName + << " (failures=" << bestHealth.failureCount + << ", avgQueueTime=" << bestHealth.avgQueueTimeMs << " ms)"; + } + return bestServerName; +} + +void TritonService::startFallbackServer() { + // Idempotent: do nothing if already running or disabled + if (!fallbackOpts_.enable || startedFallback_) return; //include fallback server in set @@ -262,10 +416,13 @@ void TritonService::preBeginJob(edm::ProcessContext const&) { std::string msg; if (verbose_) msg = "List of models for fallback server: "; - //all unserved models are provided by fallback server + // Provide all declared models with known paths via the fallback server auto& server(servers_.find(Server::fallbackName)->second); - for (const auto& [modelName, model] : unservedModels_) { - auto& modelInfo(models_.emplace(modelName, model).first->second); + for (const auto& [modelName, model] : models_) { + // Only seed models for which we have a repository path + if (model.path.empty()) + continue; + auto& modelInfo(models_.find(modelName)->second); modelInfo.servers.insert(Server::fallbackName); server.models.insert(modelName); if (verbose_) @@ -288,7 +445,9 @@ void TritonService::preBeginJob(edm::ProcessContext const&) { fallbackOpts_.command += " -r " + std::to_string(fallbackOpts_.retries); if (fallbackOpts_.wait >= 0) fallbackOpts_.command += " -w " + std::to_string(fallbackOpts_.wait); - for (const auto& [modelName, model] : unservedModels_) { + for (const auto& [modelName, model] : models_) { + if (model.path.empty()) + continue; fallbackOpts_.command += " -m " + model.path; } std::string thread_string = " -I " + std::to_string(numberOfThreads_); @@ -297,8 +456,7 @@ void TritonService::preBeginJob(edm::ProcessContext const&) { fallbackOpts_.command += " -i " + fallbackOpts_.imageName; if (!fallbackOpts_.sandboxDir.empty()) fallbackOpts_.command += " -s " + fallbackOpts_.sandboxDir; - //don't need this anymore - unservedModels_.clear(); + // models_ remains for runtime queries; nothing to clear here //get a random temporary directory if none specified if (fallbackOpts_.tempDir.empty()) { @@ -360,6 +518,25 @@ void TritonService::preBeginJob(edm::ProcessContext const&) { << output; } +void TritonService::preBeginJob(edm::ProcessContext const&) { + // Capture unassigned models *before* startFallbackServer() is called. + // startFallbackServer() seeds all known-path models into the fallback server + // set, which would make unassignedModels() return empty afterward. + const auto& unassigned = unassignedModels(); + + // Always start the fallback server so it is ready for on-demand model + // loading during retries, even when every model has a primary server. + startFallbackServer(); + + if (!unassigned.empty() && startedFallback_) { + auto& server(servers_.find(Server::fallbackName)->second); + for (const auto& modelName : unassigned) { + server.models.insert(modelName); + loadModel(modelName); + } + } +} + void TritonService::notifyCallStatus(bool status) const { if (status) --callFails_; @@ -453,3 +630,107 @@ void TritonService::fillDescriptions(edm::ConfigurationDescriptions& description descriptions.addWithDefaultLabel(desc); } + +bool TritonService::loadModel(const std::string& modelName) { + std::lock_guard lock(modelLoadMutex_); + + // Get model from models_ map (should exist from addModel during module construction) + auto mit = models_.find(modelName); + if (mit == models_.end()) { + edm::LogWarning("TritonService") << "loadModel: Model " << modelName << " not found in models_ map"; + return false; + } + + return loadModel(modelName, mit->second); +} + +bool TritonService::loadModel(const std::string& modelName, Model& model) { + // if already loaded, bump refcount + if (model.refCount > 0) { + ++model.refCount; + if (verbose_) + edm::LogInfo("TritonService") << "Model " << modelName << " already loaded, ref count: " << model.refCount; + return true; + } + + if (!startedFallback_) { + throw cms::Exception("TritonService") + << "loadModel: fallback server not started; cannot load model '" << modelName << "'"; + } + + auto sit = servers_.find(Server::fallbackName); + if (sit == servers_.end()) { + throw cms::Exception("TritonService") << "loadModel: fallback server not found"; + } + + std::unique_ptr client; + TRITON_THROW_IF_ERROR(tc::InferenceServerGrpcClient::Create( + &client, sit->second.url, false, sit->second.useSsl, sit->second.sslOptions), + "loadModel: unable to create client for fallback server"); + + TRITON_THROW_IF_ERROR(client->LoadModel(modelName), + "loadModel: failed to load model " + modelName + " on fallback server"); + + // Update state and tracking + model.refCount = 1; + model.servers.insert(Server::fallbackName); + sit->second.models.insert(modelName); + fallbackLoadedModels_.insert(modelName); + + if (verbose_) + edm::LogInfo("TritonService") << "Successfully loaded model " << modelName << " on fallback server"; + return true; +} + +bool TritonService::unloadModel(const std::string& modelName) { + std::lock_guard lock(modelLoadMutex_); + + // Get model from models_ map + auto mit = models_.find(modelName); + if (mit == models_.end()) { + edm::LogWarning("TritonService") << "unloadModel: Model " << modelName << " not found in models_ map"; + return false; + } + + return unloadModel(modelName, mit->second); +} + +bool TritonService::unloadModel(const std::string& modelName, Model& model) { + if (model.refCount == 0) { + edm::LogWarning("TritonService") << "unloadModel: Model " << modelName << " is not loaded"; + return false; + } + + if (model.refCount > 1) { + --model.refCount; + if (verbose_) + edm::LogInfo("TritonService") << "Model " << modelName << " still in use, ref count: " << model.refCount; + return true; + } + + auto sit = servers_.find(Server::fallbackName); + if (sit == servers_.end()) { + edm::LogWarning("TritonService") << "unloadModel: Fallback server not found"; + return false; + } + + if (verbose_) + edm::LogInfo("TritonService") << "Model " << modelName << " ref count is 1, unloading from fallback server"; + + std::unique_ptr client; + TRITON_THROW_IF_ERROR(tc::InferenceServerGrpcClient::Create( + &client, sit->second.url, false, sit->second.useSsl, sit->second.sslOptions), + "unloadModel: unable to create client for fallback server"); + + TRITON_THROW_IF_ERROR(client->UnloadModel(modelName), + "unloadModel: failed to unload model " + modelName + " from fallback server"); + + model.refCount = 0; + model.servers.erase(Server::fallbackName); + sit->second.models.erase(modelName); + fallbackLoadedModels_.erase(modelName); + + if (verbose_) + edm::LogInfo("TritonService") << "Successfully unloaded model " << modelName << " from fallback server"; + return true; +} diff --git a/HeterogeneousCore/SonicTriton/test/BuildFile.xml b/HeterogeneousCore/SonicTriton/test/BuildFile.xml index e4ff7a0bb56f3..ec32b11b88970 100644 --- a/HeterogeneousCore/SonicTriton/test/BuildFile.xml +++ b/HeterogeneousCore/SonicTriton/test/BuildFile.xml @@ -1,11 +1,31 @@ + + - + + + + + + + + + + + + + + + + + + + diff --git a/HeterogeneousCore/SonicTriton/test/DynamicModelLoadingProducer.cc b/HeterogeneousCore/SonicTriton/test/DynamicModelLoadingProducer.cc new file mode 100644 index 0000000000000..eb885ea29dec1 --- /dev/null +++ b/HeterogeneousCore/SonicTriton/test/DynamicModelLoadingProducer.cc @@ -0,0 +1,83 @@ +#include "HeterogeneousCore/SonicTriton/interface/TritonEDProducer.h" +#include "HeterogeneousCore/SonicTriton/interface/TritonService.h" +#include "DataFormats/TestObjects/interface/ToyProducts.h" +#include "FWCore/Framework/interface/MakerMacros.h" +#include "FWCore/MessageLogger/interface/MessageLogger.h" +#include "FWCore/ServiceRegistry/interface/Service.h" + +#include +#include +#include +#include + +// Test module that explicitly exercises dynamic model loading +// This tests the reference counting and thread safety of loadModel/unloadModel +class DynamicModelLoadingProducer : public TritonEDProducer<> { +public: + explicit DynamicModelLoadingProducer(edm::ParameterSet const& cfg) + : TritonEDProducer<>(cfg), + loadUnloadCycles_(cfg.getParameter("loadUnloadCycles")), + testConcurrency_(cfg.getParameter("testConcurrency")) { + putToken_ = produces(); + } + + void acquire(edm::Event const& iEvent, edm::EventSetup const& iSetup, Input& iInput) override { + edm::Service ts; + const std::string& modelName = client_->modelName(); + + // Test dynamic loading and unloading + if (testConcurrency_) { + // Stress test with multiple rapid load/unload cycles + for (int i = 0; i < loadUnloadCycles_; ++i) { + bool loadResult = ts->loadModel(modelName); + edm::LogInfo("DynamicModelLoadingProducer") + << "Load attempt " << i << ": " << (loadResult ? "success" : "failed"); + + // Small delay to allow other threads to interleave + if (i % 5 == 0) { + std::this_thread::yield(); + } + + bool unloadResult = ts->unloadModel(modelName); + edm::LogInfo("DynamicModelLoadingProducer") + << "Unload attempt " << i << ": " << (unloadResult ? "success" : "failed"); + } + } else { + // Simple test: load once, unload once + bool loadResult = ts->loadModel(modelName); + edm::LogInfo("DynamicModelLoadingProducer") << "Single load: " << (loadResult ? "success" : "failed"); + + bool unloadResult = ts->unloadModel(modelName); + edm::LogInfo("DynamicModelLoadingProducer") << "Single unload: " << (unloadResult ? "success" : "failed"); + } + + // Fill dummy input - use actual input from the model (gat_test expects "x" input) + // This is just to satisfy the base class requirements, not for actual inference + auto& input_x = iInput.at("x"); + auto data_x = input_x.allocate(); + // Minimal dummy data + (*data_x)[0] = std::vector{1.0f}; + input_x.setShape(0, 1, 0); + input_x.toServer(data_x); + } + + void produce(edm::Event& iEvent, edm::EventSetup const& iSetup, Output const& iOutput) override { + // Produce dummy output + iEvent.emplace(putToken_, loadUnloadCycles_); + } + + static void fillDescriptions(edm::ConfigurationDescriptions& descriptions) { + edm::ParameterSetDescription desc; + TritonClient::fillPSetDescription(desc); + desc.add("loadUnloadCycles", 1); + desc.add("testConcurrency", false); + descriptions.addWithDefaultLabel(desc); + } + +private: + int loadUnloadCycles_; + bool testConcurrency_; + edm::EDPutTokenT putToken_; +}; + +DEFINE_FWK_MODULE(DynamicModelLoadingProducer); diff --git a/HeterogeneousCore/SonicTriton/test/RefCount.cc b/HeterogeneousCore/SonicTriton/test/RefCount.cc new file mode 100644 index 0000000000000..735b334542e1a --- /dev/null +++ b/HeterogeneousCore/SonicTriton/test/RefCount.cc @@ -0,0 +1,199 @@ +#define CATCH_CONFIG_MAIN +#include "catch2/catch_all.hpp" + +#include +#include +#include + +// Standalone refcount logic test +// This tests the refcount algorithm without requiring the full TritonService infrastructure + +// Simplified model state for testing refcount logic +struct TestModelState { + std::string modelName; + std::string path; + int refCount{0}; + bool isLoaded() const { return refCount > 0; } +}; + +// Mock class that implements the same refcount logic as TritonService +class RefCountManager { +public: + // Tracks actual server load/unload calls + int serverLoadCalls{0}; + int serverUnloadCalls{0}; + + // Simulates loadModel behavior + bool loadModel(const std::string& modelName, const std::string& path = "") { + std::lock_guard lock(mutex_); + + auto& state = models_[modelName]; + if (state.modelName.empty()) + state.modelName = modelName; + if (state.path.empty() && !path.empty()) + state.path = path; + + // If already loaded, just bump refcount (no server call) + if (state.refCount > 0) { + ++state.refCount; + refCounts_[modelName] = state.refCount; + return true; + } + + // Actually "load" on server (simulated) + ++serverLoadCalls; + state.refCount = 1; + refCounts_[modelName] = state.refCount; + return true; + } + + // Simulates unloadModel behavior + bool unloadModel(const std::string& modelName) { + std::lock_guard lock(mutex_); + + auto it = models_.find(modelName); + if (it == models_.end() || it->second.refCount == 0) { + return false; // Not loaded + } + + auto& state = it->second; + + // If refcount > 1, just decrement (no server call) + if (state.refCount > 1) { + --(state.refCount); + refCounts_[modelName] = state.refCount; + return true; + } + + // Actually "unload" from server (simulated) + ++serverUnloadCalls; + refCounts_.erase(modelName); + state.refCount = 0; + return true; + } + + int getRefCount(const std::string& modelName) const { + auto it = refCounts_.find(modelName); + return (it != refCounts_.end()) ? it->second : 0; + } + +private: + std::unordered_map models_; + std::unordered_map refCounts_; + std::mutex mutex_; +}; + +TEST_CASE("RefCount: single load increments to 1", "[RefCount]") { + RefCountManager mgr; + + REQUIRE(mgr.loadModel("model_a", "/path/to/model_a")); + REQUIRE(mgr.getRefCount("model_a") == 1); + REQUIRE(mgr.serverLoadCalls == 1); +} + +TEST_CASE("RefCount: multiple loads increment without server calls", "[RefCount]") { + RefCountManager mgr; + + // First load - should call server + REQUIRE(mgr.loadModel("model_a")); + REQUIRE(mgr.getRefCount("model_a") == 1); + REQUIRE(mgr.serverLoadCalls == 1); + + // Second load - should NOT call server, just increment + REQUIRE(mgr.loadModel("model_a")); + REQUIRE(mgr.getRefCount("model_a") == 2); + REQUIRE(mgr.serverLoadCalls == 1); // Still 1 + + // Third load - should NOT call server, just increment + REQUIRE(mgr.loadModel("model_a")); + REQUIRE(mgr.getRefCount("model_a") == 3); + REQUIRE(mgr.serverLoadCalls == 1); // Still 1 +} + +TEST_CASE("RefCount: unload decrements without server call until zero", "[RefCount]") { + RefCountManager mgr; + + // Load 3 times + mgr.loadModel("model_a"); + mgr.loadModel("model_a"); + mgr.loadModel("model_a"); + REQUIRE(mgr.getRefCount("model_a") == 3); + REQUIRE(mgr.serverLoadCalls == 1); + + // First unload - decrement only, no server call + REQUIRE(mgr.unloadModel("model_a")); + REQUIRE(mgr.getRefCount("model_a") == 2); + REQUIRE(mgr.serverUnloadCalls == 0); + + // Second unload - decrement only, no server call + REQUIRE(mgr.unloadModel("model_a")); + REQUIRE(mgr.getRefCount("model_a") == 1); + REQUIRE(mgr.serverUnloadCalls == 0); + + // Third unload - should call server (refcount reaches 0) + REQUIRE(mgr.unloadModel("model_a")); + REQUIRE(mgr.getRefCount("model_a") == 0); + REQUIRE(mgr.serverUnloadCalls == 1); +} + +TEST_CASE("RefCount: unload on non-loaded model returns false", "[RefCount]") { + RefCountManager mgr; + + // Unload without loading first + REQUIRE_FALSE(mgr.unloadModel("model_a")); + REQUIRE(mgr.serverUnloadCalls == 0); +} + +TEST_CASE("RefCount: reload after full unload triggers new server load", "[RefCount]") { + RefCountManager mgr; + + // Load and fully unload + mgr.loadModel("model_a"); + mgr.unloadModel("model_a"); + REQUIRE(mgr.getRefCount("model_a") == 0); + REQUIRE(mgr.serverLoadCalls == 1); + REQUIRE(mgr.serverUnloadCalls == 1); + + // Reload - should call server again + REQUIRE(mgr.loadModel("model_a")); + REQUIRE(mgr.getRefCount("model_a") == 1); + REQUIRE(mgr.serverLoadCalls == 2); // Now 2 +} + +TEST_CASE("RefCount: multiple models are independent", "[RefCount]") { + RefCountManager mgr; + + // Load two different models + mgr.loadModel("model_a"); + mgr.loadModel("model_b"); + REQUIRE(mgr.getRefCount("model_a") == 1); + REQUIRE(mgr.getRefCount("model_b") == 1); + REQUIRE(mgr.serverLoadCalls == 2); + + // Load model_a again + mgr.loadModel("model_a"); + REQUIRE(mgr.getRefCount("model_a") == 2); + REQUIRE(mgr.getRefCount("model_b") == 1); + REQUIRE(mgr.serverLoadCalls == 2); // No new server call + + // Unload model_b completely + mgr.unloadModel("model_b"); + REQUIRE(mgr.getRefCount("model_a") == 2); + REQUIRE(mgr.getRefCount("model_b") == 0); + REQUIRE(mgr.serverUnloadCalls == 1); + + // model_a still loaded + REQUIRE(mgr.getRefCount("model_a") == 2); +} + +TEST_CASE("RefCount: path is preserved from first load", "[RefCount]") { + RefCountManager mgr; + + // First load with path + mgr.loadModel("model_a", "/path/to/model"); + REQUIRE(mgr.getRefCount("model_a") == 1); + + // Second load without path - should still work + mgr.loadModel("model_a"); + REQUIRE(mgr.getRefCount("model_a") == 2); +} diff --git a/HeterogeneousCore/SonicTriton/test/RetryActionDiffServer.cc b/HeterogeneousCore/SonicTriton/test/RetryActionDiffServer.cc new file mode 100644 index 0000000000000..9952218a82fc3 --- /dev/null +++ b/HeterogeneousCore/SonicTriton/test/RetryActionDiffServer.cc @@ -0,0 +1,71 @@ +#define CATCH_CONFIG_MAIN +#include "catch2/catch_all.hpp" + +#include "HeterogeneousCore/SonicTriton/interface/RetryActionDiffServer.h" +#include "HeterogeneousCore/SonicTriton/interface/TritonClient.h" +#include "HeterogeneousCore/SonicTriton/interface/TritonService.h" +#include "HeterogeneousCore/SonicCore/interface/RetryActionBase.h" + +#include "FWCore/ParameterSet/interface/ParameterSet.h" + +#include + +// Test double for TritonClient to observe updateServer calls without framework/services +class TestTritonClient : public TritonClient { +public: + TestTritonClient() : TritonClient() {} + + void connectToServer(const std::string& url) override { lastConnectedUrl = url; } + + void updateServer(const std::string& serverName) override { lastUpdatedServerName = serverName; } + + const std::string& lastUrl() const { return lastConnectedUrl; } + const std::string& lastServerName() const { return lastUpdatedServerName; } + +protected: + void evaluate() override {} + +private: + std::string lastConnectedUrl; + std::string lastUpdatedServerName; +}; + +TEST_CASE("RetryActionDiffServer switches to fallback via updateServer", "[RetryActionDiffServer]") { + edm::ParameterSet empty; + TestTritonClient client; + + RetryActionDiffServer action(empty, static_cast(&client)); + + // start should arm the action + action.start(); + REQUIRE(action.shouldRetry()); + + // retry should call updateServer with fallback name then disarm + action.retry(); + REQUIRE(client.lastServerName() == TritonService::Server::fallbackName); + + // second retry without re-arming should be a no-op: lastServerName unchanged + std::string afterFirst = client.lastServerName(); + action.retry(); + REQUIRE(client.lastServerName() == afterFirst); +} + +// A client that throws during updateServer to exercise error handling path +class ThrowingTritonClient : public TritonClient { +public: + ThrowingTritonClient() : TritonClient() {} + void updateServer(const std::string&) override { throw TritonException("updateServer failure"); } + +protected: + void evaluate() override {} +}; + +TEST_CASE("RetryActionDiffServer catches exceptions from updateServer", "[RetryActionDiffServer]") { + edm::ParameterSet empty; + ThrowingTritonClient client; + RetryActionDiffServer action(empty, static_cast(&client)); + action.start(); + + // Should not throw despite client throwing internally; action disarms afterward + REQUIRE_NOTHROW(action.retry()); +} diff --git a/HeterogeneousCore/SonicTriton/test/tritonTest_cfg.py b/HeterogeneousCore/SonicTriton/test/tritonTest_cfg.py index 33d6a9c60aad4..7a5e88a4d6d6a 100644 --- a/HeterogeneousCore/SonicTriton/test/tritonTest_cfg.py +++ b/HeterogeneousCore/SonicTriton/test/tritonTest_cfg.py @@ -9,6 +9,7 @@ "TritonGraphFilter": ["gat_test"], "TritonGraphAnalyzer": ["gat_test"], "TritonIdentityProducer": ["ragged_io"], + "DynamicModelLoadingProducer": ["gat_test"], } # other choices @@ -21,6 +22,9 @@ parser.add_argument("--brief", default=False, action="store_true", help="briefer output for graph modules") parser.add_argument("--unittest", default=False, action="store_true", help="unit test mode: reduce input sizes") parser.add_argument("--testother", default=False, action="store_true", help="also test gRPC communication if shared memory enabled, or vice versa") +parser.add_argument("--loadUnloadCycles", default=3, type=int, help="number of load/unload cycles for dynamic model loading test") +parser.add_argument("--testConcurrency", default=False, action="store_true", help="enable concurrent stress test for dynamic model loading") +options = parser.parse_args() options = getOptions(parser, verbose=True) @@ -83,6 +87,10 @@ processModule.edgeMin = cms.uint32(8000) processModule.edgeMax = cms.uint32(15000) processModule.brief = cms.bool(options.brief) + elif module=="DynamicModelLoadingProducer": + # Configure dynamic model loading test (requires explicit model control mode, enabled by default in cmsTriton) + processModule.loadUnloadCycles = cms.int32(options.loadUnloadCycles) + processModule.testConcurrency = cms.bool(options.testConcurrency) process.p += processModule if options.testother: # clone modules to test both gRPC and shared memory From 53716fa35a390f777b44d826d1e86c7a8f6ebece Mon Sep 17 00:00:00 2001 From: Martin Date: Tue, 14 Jul 2026 09:21:47 -0500 Subject: [PATCH 2/9] PR comments edit1 --- .../SonicCore/src/SonicClientBase.cc | 1 - .../SonicTriton/interface/TritonClient.h | 1 - .../SonicTriton/interface/TritonService.h | 8 - .../SonicTriton/python/customize.py | 36 ++-- .../SonicTriton/src/RetryActionDiffServer.cc | 5 +- .../SonicTriton/src/TritonClient.cc | 28 +-- .../SonicTriton/src/TritonService.cc | 14 +- .../SonicTriton/test/BuildFile.xml | 6 +- .../SonicTriton/test/RefCount.cc | 199 ------------------ ...erver.cc => test_RetryActionDiffServer.cc} | 4 - 10 files changed, 28 insertions(+), 274 deletions(-) delete mode 100644 HeterogeneousCore/SonicTriton/test/RefCount.cc rename HeterogeneousCore/SonicTriton/test/{RetryActionDiffServer.cc => test_RetryActionDiffServer.cc} (92%) diff --git a/HeterogeneousCore/SonicCore/src/SonicClientBase.cc b/HeterogeneousCore/SonicCore/src/SonicClientBase.cc index 6ed10089e1bd3..05cee50a2b5f0 100644 --- a/HeterogeneousCore/SonicCore/src/SonicClientBase.cc +++ b/HeterogeneousCore/SonicCore/src/SonicClientBase.cc @@ -74,7 +74,6 @@ void SonicClientBase::finish(bool success, std::exception_ptr eptr) { edm::LogInfo("SonicClientBase") << "Calling retry()"; // retry() must trigger eval() or finish() action->retry(); - // return because another finish() was already called inside client->evaluate() return; } } diff --git a/HeterogeneousCore/SonicTriton/interface/TritonClient.h b/HeterogeneousCore/SonicTriton/interface/TritonClient.h index 74d6649f6c5e8..15692346d012f 100644 --- a/HeterogeneousCore/SonicTriton/interface/TritonClient.h +++ b/HeterogeneousCore/SonicTriton/interface/TritonClient.h @@ -56,7 +56,6 @@ class TritonClient : public SonicClient { const TritonService* localService() const; std::string modelName() const { return options_[0].model_name_; } std::string serverName() const { return serverName_; } - virtual void connectToServer(const std::string& url); virtual void updateServer(const std::string& serverName); virtual void switchToFallback(); diff --git a/HeterogeneousCore/SonicTriton/interface/TritonService.h b/HeterogeneousCore/SonicTriton/interface/TritonService.h index 9be006dfab615..3be31cb9f2ecd 100644 --- a/HeterogeneousCore/SonicTriton/interface/TritonService.h +++ b/HeterogeneousCore/SonicTriton/interface/TritonService.h @@ -136,14 +136,6 @@ class TritonService { // return the best server for retry, ignore the current server std::optional getBestServer(const std::string& modelName, const std::string& IgnoreServer = "") const; - // helper functions to get server statistics? - // - getServerSideStatus() - // - updateServerStatus() - // - loop over servers_ get statistics - // - getBestServer(model) - // - call updateServerStatus() - // - loop over servers_ get their statistics, compute metric, return server name - const std::string& pid() const { return pid_; } void notifyCallStatus(bool status) const; diff --git a/HeterogeneousCore/SonicTriton/python/customize.py b/HeterogeneousCore/SonicTriton/python/customize.py index e633d169e919e..0f397cea00c3f 100644 --- a/HeterogeneousCore/SonicTriton/python/customize.py +++ b/HeterogeneousCore/SonicTriton/python/customize.py @@ -15,10 +15,10 @@ def getParser(): parser.add_argument("--maxEvents", default=-1, type=int, help="Number of events to process (-1 for all)") parser.add_argument("--address", nargs=3, action="append", metavar=("NAME", "HOST", "PORT"), dest="addresses", default=[], - help="Triton server entry: name host port (repeatable, e.g. --address server1 0.0.0.0 8011)") + help="Triton server entry: name host port (repeatable, e.g. --address server1 0.0.0.0 8011 --address server2 0.0.0.0 8021)") parser.add_argument("--timeout", default=30, type=int, help="timeout for requests") parser.add_argument("--timeoutUnit", default="seconds", type=str, help="unit for timeout") - parser.add_argument("--params", default="", type=str, help="json file containing server address/port(single-server)") + parser.add_argument("--params", default="", type=str, help="json file containing server address/port(s) (single server dict, or list of server dicts[{'name': ..., 'address': ..., 'port': ...}])") parser.add_argument("--threads", default=1, type=int, help="number of threads") parser.add_argument("--streams", default=0, type=int, help="number of streams") parser.add_argument("--verbose", default=False, action="store_true", help="enable all verbose output") @@ -43,26 +43,22 @@ def getParser(): def getOptions(parser, verbose=False): options = parser.parse_args() - # Legacy --params support: loads a single server and appends it to options.addresses +def getOptions(parser, verbose=False): + options = parser.parse_args() + if len(options.params) > 0: with open(options.params, 'r') as pfile: pdict = json.load(pfile) - name = pdict.get("name", "default") - host = pdict["address"] - port = str(int(pdict["port"])) - options.addresses.append([name, host, port]) - if verbose: - print("server (from params) = {}:{} [{}]".format(host, port, name)) - - if verbose: - for name, host, port in options.addresses: - print("server = {}:{} [{}]".format(host, port, name)) - if len(options.params)>0: - with open(options.params,'r') as pfile: - pdict = json.load(pfile) - options.address = pdict["address"] - options.port = int(pdict["port"]) - if verbose: print("server = "+options.address+":"+str(options.port)) + + server_list = pdict if isinstance(pdict, list) else [pdict] + + for entry in server_list: + name = entry.get("name", "default") + host = entry["address"] + port = str(int(entry["port"])) + options.addresses.append([name, host, port]) + if verbose: + print("server (from params) = {}:{} [{}]".format(host, port, name)) return options @@ -92,6 +88,8 @@ def applyOptions(process, options, applyToModules=False): if len(options.fallbackName)>0: process.TritonService.fallback.instanceBaseName = options.fallbackName for name, host, port in options.addresses: + if options.verbose: + print("server = {}:{} [{}]".format(host, port, name)) process.TritonService.servers.append( dict( name = name, diff --git a/HeterogeneousCore/SonicTriton/src/RetryActionDiffServer.cc b/HeterogeneousCore/SonicTriton/src/RetryActionDiffServer.cc index 8a096f1bfd672..f803a3c7a5b09 100644 --- a/HeterogeneousCore/SonicTriton/src/RetryActionDiffServer.cc +++ b/HeterogeneousCore/SonicTriton/src/RetryActionDiffServer.cc @@ -27,14 +27,13 @@ void RetryActionDiffServer::retry() { auto bestServerName = ts->getBestServer(tritonClient->modelName(), tritonClient->serverName()); if (bestServerName) { - edm::LogInfo("RetryActionDiffServer") << "Got best server from service "; + edm::LogInfo("RetryActionDiffServer") << "Got best server from service"; tritonClient->updateServer(*bestServerName); edm::LogInfo("RetryActionDiffServer") << "eval() with new server"; eval(); return; } else { - edm::LogWarning("RetryActionDiffServer") - << "No alternative server found for model " << tritonClient->modelName() << ". Now call client->finish()"; + edm::LogWarning("RetryActionDiffServer") << "No alternative server found for model " << tritonClient->modelName(); finish(false); return; } diff --git a/HeterogeneousCore/SonicTriton/src/TritonClient.cc b/HeterogeneousCore/SonicTriton/src/TritonClient.cc index bc223440a68b4..a4781d148c16b 100644 --- a/HeterogeneousCore/SonicTriton/src/TritonClient.cc +++ b/HeterogeneousCore/SonicTriton/src/TritonClient.cc @@ -441,7 +441,6 @@ void TritonClient::evaluate() { element.second.prepare(); } }); - //edm::LogInfo("TritonClient") << "evaluate() return 1"; if (!success) return; @@ -451,7 +450,6 @@ void TritonClient::evaluate() { if (verbose()) start_status = getServerSideStatus(); }); - //edm::LogInfo("TritonClient") << "evaluate() return 2"; if (!success) return; @@ -459,11 +457,11 @@ void TritonClient::evaluate() { // Reset before each inference attempt: false=retryable/not-yet-succeeded. // The callback sets this to true on success before calling doneWaiting(nullptr). // If the launch throws (lambda destroyed inside AsyncInferMulti), the holder destructor - // calls doneWaiting(nullptr) with inferSuccess_=false → finish(false) [retryable] on TBB. + // calls doneWaiting(nullptr) with inferSuccess_=false -> finish(false) [retryable] on TBB. inferSuccess_ = false; // Create holder[finish]: when doneWaiting() is called, finish() is scheduled as a TBB task - // so retry logic (finish→retry→updateServer) never runs on the gRPC callback thread. + // so retry logic (finish->retry->updateServer) never runs on the gRPC callback thread. auto* finishTask = edm::make_waiting_task([this](std::exception_ptr const* excptr) { if (excptr) finish(false, *excptr); // non-retryable exception @@ -515,10 +513,9 @@ void TritonClient::evaluate() { localService()); } catch (TritonException& e) { e.convertToWarning(); - // fh destructor already called doneWaiting(nullptr) → finish(false) scheduled on TBB + // fh destructor already called doneWaiting(nullptr) -> finish(false) scheduled on TBB } } else { - //edm::LogInfo("TritonClient") << "evaluate() return 4"; //blocking call std::vector resultsTmp; success = handle_exception([&]() { @@ -633,9 +630,10 @@ void TritonClient::updateServer(const std::string& serverName) { serverType_ = server.type; edm::LogInfo("TritonDiscovery") << debugName_ << " assigned server: " << server.url; //enforce sync mode for fallback CPU server to avoid contention - //todo: could enforce async mode otherwise (unless mode was specified by user?) if (serverType_ == TritonServerType::LocalCPU) setMode(SonicMode::Sync); + if (serverType_ == TritonServerType::Remote) + setMode(SonicMode::Async); isLocal_ = serverType_ == TritonServerType::LocalCPU or serverType_ == TritonServerType::LocalGPU; // updateServer() is always called from a TBB thread (via finish() -> retry() -> updateServer()), @@ -685,21 +683,5 @@ void TritonClient::fillPSetDescription(edm::ParameterSetDescription& iDesc) { iDesc.add("Client", descClient); } -void TritonClient::connectToServer(const std::string& url) { - // Update client state for a generic remote server - serverType_ = TritonServerType::Remote; - isLocal_ = false; - - edm::LogInfo("TritonDiscovery") << debugName_ << " connecting to server: " << url; - - // Use default SSL options - triton::client::SslOptions sslOptions; - bool useSsl = false; // Assuming no SSL for direct URL connection - - // Connect to the server - TRITON_THROW_IF_ERROR(triton::client::InferenceServerGrpcClient::Create(&client_, url, false, useSsl, sslOptions), - "TritonClient::connectToServer(): unable to create inference context"); -} - //constructor for testing TritonClient::TritonClient() : SonicClient(makeMinimalSonicParamsForTest(), "TritonClient_test", "TritonClient") {} diff --git a/HeterogeneousCore/SonicTriton/src/TritonService.cc b/HeterogeneousCore/SonicTriton/src/TritonService.cc index b10ffacadcaa5..5f4633717cfcc 100644 --- a/HeterogeneousCore/SonicTriton/src/TritonService.cc +++ b/HeterogeneousCore/SonicTriton/src/TritonService.cc @@ -327,16 +327,8 @@ void TritonService::updateServerHealth(const std::string& modelName) const { health.avgQueueTimeMs = (queue_count > 0) ? avgQueueTimeMs / queue_count : 0.0; health.avgInferTimeMs = (infer_count > 0) ? avgInferTimeMs / infer_count : 0.0; - } catch (const TritonException& e) { - // mark existing entry unhealthy if present - tbb::concurrent_hash_map::accessor acc; - if (serversHealth_.find(acc, serverName)) { - ServerHealth& health = acc->second; - health.live = false; - health.ready = false; - } } catch (const std::exception& e) { - // fallback for other exceptions + // mark existing entry unhealthy if present tbb::concurrent_hash_map::accessor acc; if (serversHealth_.find(acc, serverName)) { ServerHealth& health = acc->second; @@ -348,7 +340,7 @@ void TritonService::updateServerHealth(const std::string& modelName) const { } std::optional TritonService::getBestServer(const std::string& modelName, - const std::string& IgnoreServer) const { + const std::string& ignoreServer) const { std::optional bestServerName; ServerHealth bestHealth; @@ -357,7 +349,7 @@ std::optional TritonService::getBestServer(const std::string& model edm::LogInfo("TritonService") << "Getting best server"; for (auto& [serverName, server] : servers_) { - if (serverName == IgnoreServer) { + if (serverName == ignoreServer) { edm::LogInfo("TritonService") << serverName << " is ignored"; continue; // skip ignored server } diff --git a/HeterogeneousCore/SonicTriton/test/BuildFile.xml b/HeterogeneousCore/SonicTriton/test/BuildFile.xml index ec32b11b88970..4e8b765205e6f 100644 --- a/HeterogeneousCore/SonicTriton/test/BuildFile.xml +++ b/HeterogeneousCore/SonicTriton/test/BuildFile.xml @@ -14,16 +14,12 @@ - + - - - - diff --git a/HeterogeneousCore/SonicTriton/test/RefCount.cc b/HeterogeneousCore/SonicTriton/test/RefCount.cc deleted file mode 100644 index 735b334542e1a..0000000000000 --- a/HeterogeneousCore/SonicTriton/test/RefCount.cc +++ /dev/null @@ -1,199 +0,0 @@ -#define CATCH_CONFIG_MAIN -#include "catch2/catch_all.hpp" - -#include -#include -#include - -// Standalone refcount logic test -// This tests the refcount algorithm without requiring the full TritonService infrastructure - -// Simplified model state for testing refcount logic -struct TestModelState { - std::string modelName; - std::string path; - int refCount{0}; - bool isLoaded() const { return refCount > 0; } -}; - -// Mock class that implements the same refcount logic as TritonService -class RefCountManager { -public: - // Tracks actual server load/unload calls - int serverLoadCalls{0}; - int serverUnloadCalls{0}; - - // Simulates loadModel behavior - bool loadModel(const std::string& modelName, const std::string& path = "") { - std::lock_guard lock(mutex_); - - auto& state = models_[modelName]; - if (state.modelName.empty()) - state.modelName = modelName; - if (state.path.empty() && !path.empty()) - state.path = path; - - // If already loaded, just bump refcount (no server call) - if (state.refCount > 0) { - ++state.refCount; - refCounts_[modelName] = state.refCount; - return true; - } - - // Actually "load" on server (simulated) - ++serverLoadCalls; - state.refCount = 1; - refCounts_[modelName] = state.refCount; - return true; - } - - // Simulates unloadModel behavior - bool unloadModel(const std::string& modelName) { - std::lock_guard lock(mutex_); - - auto it = models_.find(modelName); - if (it == models_.end() || it->second.refCount == 0) { - return false; // Not loaded - } - - auto& state = it->second; - - // If refcount > 1, just decrement (no server call) - if (state.refCount > 1) { - --(state.refCount); - refCounts_[modelName] = state.refCount; - return true; - } - - // Actually "unload" from server (simulated) - ++serverUnloadCalls; - refCounts_.erase(modelName); - state.refCount = 0; - return true; - } - - int getRefCount(const std::string& modelName) const { - auto it = refCounts_.find(modelName); - return (it != refCounts_.end()) ? it->second : 0; - } - -private: - std::unordered_map models_; - std::unordered_map refCounts_; - std::mutex mutex_; -}; - -TEST_CASE("RefCount: single load increments to 1", "[RefCount]") { - RefCountManager mgr; - - REQUIRE(mgr.loadModel("model_a", "/path/to/model_a")); - REQUIRE(mgr.getRefCount("model_a") == 1); - REQUIRE(mgr.serverLoadCalls == 1); -} - -TEST_CASE("RefCount: multiple loads increment without server calls", "[RefCount]") { - RefCountManager mgr; - - // First load - should call server - REQUIRE(mgr.loadModel("model_a")); - REQUIRE(mgr.getRefCount("model_a") == 1); - REQUIRE(mgr.serverLoadCalls == 1); - - // Second load - should NOT call server, just increment - REQUIRE(mgr.loadModel("model_a")); - REQUIRE(mgr.getRefCount("model_a") == 2); - REQUIRE(mgr.serverLoadCalls == 1); // Still 1 - - // Third load - should NOT call server, just increment - REQUIRE(mgr.loadModel("model_a")); - REQUIRE(mgr.getRefCount("model_a") == 3); - REQUIRE(mgr.serverLoadCalls == 1); // Still 1 -} - -TEST_CASE("RefCount: unload decrements without server call until zero", "[RefCount]") { - RefCountManager mgr; - - // Load 3 times - mgr.loadModel("model_a"); - mgr.loadModel("model_a"); - mgr.loadModel("model_a"); - REQUIRE(mgr.getRefCount("model_a") == 3); - REQUIRE(mgr.serverLoadCalls == 1); - - // First unload - decrement only, no server call - REQUIRE(mgr.unloadModel("model_a")); - REQUIRE(mgr.getRefCount("model_a") == 2); - REQUIRE(mgr.serverUnloadCalls == 0); - - // Second unload - decrement only, no server call - REQUIRE(mgr.unloadModel("model_a")); - REQUIRE(mgr.getRefCount("model_a") == 1); - REQUIRE(mgr.serverUnloadCalls == 0); - - // Third unload - should call server (refcount reaches 0) - REQUIRE(mgr.unloadModel("model_a")); - REQUIRE(mgr.getRefCount("model_a") == 0); - REQUIRE(mgr.serverUnloadCalls == 1); -} - -TEST_CASE("RefCount: unload on non-loaded model returns false", "[RefCount]") { - RefCountManager mgr; - - // Unload without loading first - REQUIRE_FALSE(mgr.unloadModel("model_a")); - REQUIRE(mgr.serverUnloadCalls == 0); -} - -TEST_CASE("RefCount: reload after full unload triggers new server load", "[RefCount]") { - RefCountManager mgr; - - // Load and fully unload - mgr.loadModel("model_a"); - mgr.unloadModel("model_a"); - REQUIRE(mgr.getRefCount("model_a") == 0); - REQUIRE(mgr.serverLoadCalls == 1); - REQUIRE(mgr.serverUnloadCalls == 1); - - // Reload - should call server again - REQUIRE(mgr.loadModel("model_a")); - REQUIRE(mgr.getRefCount("model_a") == 1); - REQUIRE(mgr.serverLoadCalls == 2); // Now 2 -} - -TEST_CASE("RefCount: multiple models are independent", "[RefCount]") { - RefCountManager mgr; - - // Load two different models - mgr.loadModel("model_a"); - mgr.loadModel("model_b"); - REQUIRE(mgr.getRefCount("model_a") == 1); - REQUIRE(mgr.getRefCount("model_b") == 1); - REQUIRE(mgr.serverLoadCalls == 2); - - // Load model_a again - mgr.loadModel("model_a"); - REQUIRE(mgr.getRefCount("model_a") == 2); - REQUIRE(mgr.getRefCount("model_b") == 1); - REQUIRE(mgr.serverLoadCalls == 2); // No new server call - - // Unload model_b completely - mgr.unloadModel("model_b"); - REQUIRE(mgr.getRefCount("model_a") == 2); - REQUIRE(mgr.getRefCount("model_b") == 0); - REQUIRE(mgr.serverUnloadCalls == 1); - - // model_a still loaded - REQUIRE(mgr.getRefCount("model_a") == 2); -} - -TEST_CASE("RefCount: path is preserved from first load", "[RefCount]") { - RefCountManager mgr; - - // First load with path - mgr.loadModel("model_a", "/path/to/model"); - REQUIRE(mgr.getRefCount("model_a") == 1); - - // Second load without path - should still work - mgr.loadModel("model_a"); - REQUIRE(mgr.getRefCount("model_a") == 2); -} diff --git a/HeterogeneousCore/SonicTriton/test/RetryActionDiffServer.cc b/HeterogeneousCore/SonicTriton/test/test_RetryActionDiffServer.cc similarity index 92% rename from HeterogeneousCore/SonicTriton/test/RetryActionDiffServer.cc rename to HeterogeneousCore/SonicTriton/test/test_RetryActionDiffServer.cc index 9952218a82fc3..fdb9976c1f2ac 100644 --- a/HeterogeneousCore/SonicTriton/test/RetryActionDiffServer.cc +++ b/HeterogeneousCore/SonicTriton/test/test_RetryActionDiffServer.cc @@ -15,18 +15,14 @@ class TestTritonClient : public TritonClient { public: TestTritonClient() : TritonClient() {} - void connectToServer(const std::string& url) override { lastConnectedUrl = url; } - void updateServer(const std::string& serverName) override { lastUpdatedServerName = serverName; } - const std::string& lastUrl() const { return lastConnectedUrl; } const std::string& lastServerName() const { return lastUpdatedServerName; } protected: void evaluate() override {} private: - std::string lastConnectedUrl; std::string lastUpdatedServerName; }; From f43154b92ff8764184efcb0cb6f5edf37b0782e5 Mon Sep 17 00:00:00 2001 From: Martin Date: Tue, 14 Jul 2026 11:14:54 -0500 Subject: [PATCH 3/9] Fix json support --- HeterogeneousCore/SonicTriton/python/customize.py | 1 + 1 file changed, 1 insertion(+) diff --git a/HeterogeneousCore/SonicTriton/python/customize.py b/HeterogeneousCore/SonicTriton/python/customize.py index 0f397cea00c3f..e7da1868cd2e8 100644 --- a/HeterogeneousCore/SonicTriton/python/customize.py +++ b/HeterogeneousCore/SonicTriton/python/customize.py @@ -1,4 +1,5 @@ import FWCore.ParameterSet.Config as cms +import json def getDefaultClientPSet(): from HeterogeneousCore.SonicTriton.TritonGraphAnalyzer import TritonGraphAnalyzer From b9a4f8e6f4e87d0426f71989f720adec5c425db6 Mon Sep 17 00:00:00 2001 From: Martin Date: Tue, 14 Jul 2026 13:30:17 -0500 Subject: [PATCH 4/9] add retry action test script and clean up tests --- .../SonicTriton/test/BuildFile.xml | 4 +-- .../SonicTriton/test/retry_action_same.sh | 36 +++++++++++++++++++ 2 files changed, 37 insertions(+), 3 deletions(-) create mode 100755 HeterogeneousCore/SonicTriton/test/retry_action_same.sh diff --git a/HeterogeneousCore/SonicTriton/test/BuildFile.xml b/HeterogeneousCore/SonicTriton/test/BuildFile.xml index 4e8b765205e6f..7d5ba1e0acefc 100644 --- a/HeterogeneousCore/SonicTriton/test/BuildFile.xml +++ b/HeterogeneousCore/SonicTriton/test/BuildFile.xml @@ -1,7 +1,5 @@ - - @@ -20,7 +18,7 @@ - + diff --git a/HeterogeneousCore/SonicTriton/test/retry_action_same.sh b/HeterogeneousCore/SonicTriton/test/retry_action_same.sh new file mode 100755 index 0000000000000..1a2c51ba4c4ae --- /dev/null +++ b/HeterogeneousCore/SonicTriton/test/retry_action_same.sh @@ -0,0 +1,36 @@ +#!/bin/bash + +LOCALTOP=$1 + +# Start the server +cmsTriton -v -P 8010 -n server1 -f -L -m /cvmfs/cms.cern.ch/el9_amd64_gcc12/cms/cmssw/CMSSW_15_1_0_pre6/external/el9_amd64_gcc12/data/HeterogeneousCore/SonicTriton/data/models/gat_test/config.pbtxt start & + +# Sleep to allow the server to initialize +sleep 60 + +# Get the true PID of the server +SERVER_PID=$(ps -e -o pid,cmd | grep "tritonserver" | grep -v "grep" | awk '{print $1}') + +if [ -z "$SERVER_PID" ]; then + echo "Server process could not be started." + exit 1 +fi + +echo "Server started with PID $SERVER_PID" + +ps aux | grep tritonserver + +# Start the client +cmsRun ${LOCALTOP}/src/HeterogeneousCore/SonicTriton/test/tritonTest_cfg.py --maxEvents 100 --modules TritonIdentityProducer --models ragged_io --address server1 0.0.0.0 8011 --tries 10 --verbose & + +# Allow the client some time to run +sleep 30 + +# Kill the server +echo "Killing server with PID $SERVER_PID" +kill $SERVER_PID +#echo "Server process with PID $SERVER_PID has been killed." + +# Wait for client process to complete +wait +echo "Client process completed." From 14a5e2946b3af8876f14bfc189106d193a02e03b Mon Sep 17 00:00:00 2001 From: Martin Date: Tue, 14 Jul 2026 15:06:48 -0500 Subject: [PATCH 5/9] Update allowedTries config --- HeterogeneousCore/SonicCore/README.md | 15 ++++++---- .../slimming/patElectronDRNCorrector_cfi.py | 7 ++++- .../slimming/patPhotonDRNCorrector_cfi.py | 7 ++++- .../python/pfHiggsInteractionNet_cff.py | 7 ++++- .../python/pfParticleNetAK4_cff.py | 7 ++++- .../python/pfParticleNetFromMiniAODAK4_cff.py | 28 ++++++++++++++++--- .../python/pfParticleNetFromMiniAODAK8_cff.py | 7 ++++- .../ONNXRuntime/python/pfParticleNet_cff.py | 21 ++++++++++++-- .../python/pfParticleTransformerAK4_cff.py | 7 ++++- .../pfUnifiedParticleTransformerAK4_cff.py | 7 ++++- .../SCEnergyCorrectorDRNProducer_cfi.py | 14 ++++++++-- .../RecoTau/python/tools/runTauIdMVA.py | 21 ++++++++++++-- 12 files changed, 124 insertions(+), 24 deletions(-) diff --git a/HeterogeneousCore/SonicCore/README.md b/HeterogeneousCore/SonicCore/README.md index 71d51cfb5d55f..eafcf3ca098ae 100644 --- a/HeterogeneousCore/SonicCore/README.md +++ b/HeterogeneousCore/SonicCore/README.md @@ -43,12 +43,17 @@ process.MyProducer = cms.EDProducer("MyProducer", Client = cms.PSet( # necessary client options go here mode = cms.string("Sync"), - allowedTries = cms.untracked.uint32(0), + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(0) + ) + ) ) ) ``` These parameters can be prepopulated and validated by the client using `fillDescriptions()`. -The `mode` and `allowedTries` parameters are always necessary (example values are shown here, but other values are also allowed). +The `mode` and `Retry` parameters are always necessary (example values are shown here, but other values are also allowed). These parameters are described in the next section. In addition, there is a `SonicOneEDAnalyzer` class template for user analysis, e.g. to produce simple ROOT files. @@ -110,9 +115,9 @@ For the `Sync` and `PseudoAsync` modes, `finish()` should be called at the end o For the `Async` mode, `finish()` should be called inside the communication protocol callback function (implementations may vary). When `finish()` is called, the success or failure of the call should be conveyed. -If a call fails, it can optionally be retried. This is only allowed if the call failure does not cause an exception. +If a call fails without raising an exception, it can be retried through an ordered chain of retry actions rather than a single fixed number of tries. +The chain is configured per client through a `Retry` `VPSet` parameter, where each `PSet` specifies a `retryType` plus any action-specific parameters; VPSet order is try order. Therefore, if retrying is desired, any exception should be converted to a `LogWarning` or `LogError` message by the client. -A Python configuration parameter can be provided to enable retries with a specified maximum number of allowed tries. The client must also provide a static method `fillPSetDescription()` to populate its parameters in the `fillDescriptions()` for the producers that use the client: ```cpp @@ -126,6 +131,6 @@ void MyClient::fillPSetDescription(edm::ParameterSetDescription& iDesc) { As indicated, the `fillBasePSetDescription()` function should always be applied to the `descClient` object, to ensure that it includes the necessary parameters. -(Calling `fillBasePSetDescription(descClient, false)` will omit the `allowedTries` parameter, disabling retries.) +(Calling `fillBasePSetDescription(descClient, false)` will omit the `Retry` parameter, disabling retries.) Example client code can be found in the `interface` and `src` directories of the other Sonic packages in this repository. diff --git a/PhysicsTools/PatAlgos/python/slimming/patElectronDRNCorrector_cfi.py b/PhysicsTools/PatAlgos/python/slimming/patElectronDRNCorrector_cfi.py index 94fe00cb1b93a..fa90f55215a6a 100644 --- a/PhysicsTools/PatAlgos/python/slimming/patElectronDRNCorrector_cfi.py +++ b/PhysicsTools/PatAlgos/python/slimming/patElectronDRNCorrector_cfi.py @@ -7,7 +7,12 @@ rhoName = 'fixedGridRhoFastjetAll', Client = patElectronDRNCorrectionProducer.Client.clone( mode = 'Async', - allowedTries = 1, + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(1) + ) + ), modelName = 'electronObjectEnsemble', modelConfigPath = 'RecoEgamma/EgammaElectronProducers/data/models/electronObjectEnsemble/config.pbtxt', timeout = 10 diff --git a/PhysicsTools/PatAlgos/python/slimming/patPhotonDRNCorrector_cfi.py b/PhysicsTools/PatAlgos/python/slimming/patPhotonDRNCorrector_cfi.py index cae2e38452447..4c6c6c2d23230 100644 --- a/PhysicsTools/PatAlgos/python/slimming/patPhotonDRNCorrector_cfi.py +++ b/PhysicsTools/PatAlgos/python/slimming/patPhotonDRNCorrector_cfi.py @@ -6,7 +6,12 @@ rhoName = 'fixedGridRhoFastjetAll', Client = patPhotonDRNCorrectionProducer.Client.clone( mode = 'Async', - allowedTries = 1, + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(1) + ) + ), modelName = 'photonObjectCombined', modelConfigPath = 'RecoEgamma/EgammaPhotonProducers/data/models/photonObjectCombined/config.pbtxt', timeout = 10 diff --git a/RecoBTag/ONNXRuntime/python/pfHiggsInteractionNet_cff.py b/RecoBTag/ONNXRuntime/python/pfHiggsInteractionNet_cff.py index 3195f81c51b15..cc4ef704354cc 100644 --- a/RecoBTag/ONNXRuntime/python/pfHiggsInteractionNet_cff.py +++ b/RecoBTag/ONNXRuntime/python/pfHiggsInteractionNet_cff.py @@ -33,7 +33,12 @@ modelConfigPath = cms.FileInPath("RecoBTag/Combined/data/models/higgsInteractionNet/config.pbtxt"), modelVersion = cms.string(""), verbose = cms.untracked.bool(False), - allowedTries = cms.untracked.uint32(0), + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(0) + ) + ), useSharedMemory = cms.untracked.bool(True), compression = cms.untracked.string(""), ), diff --git a/RecoBTag/ONNXRuntime/python/pfParticleNetAK4_cff.py b/RecoBTag/ONNXRuntime/python/pfParticleNetAK4_cff.py index f524eb17aee33..6ec54638152a0 100644 --- a/RecoBTag/ONNXRuntime/python/pfParticleNetAK4_cff.py +++ b/RecoBTag/ONNXRuntime/python/pfParticleNetAK4_cff.py @@ -48,7 +48,12 @@ modelConfigPath = cms.FileInPath("RecoBTag/Combined/data/models/particlenet_AK4/config.pbtxt"), modelVersion = cms.string(""), verbose = cms.untracked.bool(False), - allowedTries = cms.untracked.uint32(0), + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(0) + ) + ), useSharedMemory = cms.untracked.bool(True), compression = cms.untracked.string(""), ), diff --git a/RecoBTag/ONNXRuntime/python/pfParticleNetFromMiniAODAK4_cff.py b/RecoBTag/ONNXRuntime/python/pfParticleNetFromMiniAODAK4_cff.py index 10270a64e6a4a..ab59555102378 100644 --- a/RecoBTag/ONNXRuntime/python/pfParticleNetFromMiniAODAK4_cff.py +++ b/RecoBTag/ONNXRuntime/python/pfParticleNetFromMiniAODAK4_cff.py @@ -57,7 +57,12 @@ modelConfigPath = cms.FileInPath("RecoBTag/Combined/data/models/particleNetFromMiniAODAK4CHSCentral/config.pbtxt"), modelVersion = cms.string(""), verbose = cms.untracked.bool(False), - allowedTries = cms.untracked.uint32(0), + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(0) + ) + ), useSharedMemory = cms.untracked.bool(True), compression = cms.untracked.string(""), ), @@ -81,7 +86,12 @@ modelConfigPath = cms.FileInPath("RecoBTag/Combined/data/models/particleNetFromMiniAODAK4CHSForward/config.pbtxt"), modelVersion = cms.string(""), verbose = cms.untracked.bool(False), - allowedTries = cms.untracked.uint32(0), + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(0) + ) + ), useSharedMemory = cms.untracked.bool(True), compression = cms.untracked.string(""), ), @@ -105,7 +115,12 @@ modelConfigPath = cms.FileInPath("RecoBTag/Combined/data/models/particleNetFromMiniAODAK4PuppiCentral/config.pbtxt"), modelVersion = cms.string(""), verbose = cms.untracked.bool(False), - allowedTries = cms.untracked.uint32(0), + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(0) + ) + ), useSharedMemory = cms.untracked.bool(True), compression = cms.untracked.string(""), ), @@ -129,7 +144,12 @@ modelConfigPath = cms.FileInPath("RecoBTag/Combined/data/models/particleNetFromMiniAODAK4PuppiForward/config.pbtxt"), modelVersion = cms.string(""), verbose = cms.untracked.bool(False), - allowedTries = cms.untracked.uint32(0), + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(0) + ) + ), useSharedMemory = cms.untracked.bool(True), compression = cms.untracked.string(""), ), diff --git a/RecoBTag/ONNXRuntime/python/pfParticleNetFromMiniAODAK8_cff.py b/RecoBTag/ONNXRuntime/python/pfParticleNetFromMiniAODAK8_cff.py index 50825e9abccd6..529dfce82eb3d 100644 --- a/RecoBTag/ONNXRuntime/python/pfParticleNetFromMiniAODAK8_cff.py +++ b/RecoBTag/ONNXRuntime/python/pfParticleNetFromMiniAODAK8_cff.py @@ -33,7 +33,12 @@ modelConfigPath = cms.FileInPath("RecoBTag/Combined/data/models/particleNetFromMiniAODAK8/config.pbtxt"), modelVersion = cms.string(""), verbose = cms.untracked.bool(False), - allowedTries = cms.untracked.uint32(0), + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(0) + ) + ), useSharedMemory = cms.untracked.bool(True), compression = cms.untracked.string(""), ), diff --git a/RecoBTag/ONNXRuntime/python/pfParticleNet_cff.py b/RecoBTag/ONNXRuntime/python/pfParticleNet_cff.py index cb007ca23fc56..93df39ca9c9b3 100644 --- a/RecoBTag/ONNXRuntime/python/pfParticleNet_cff.py +++ b/RecoBTag/ONNXRuntime/python/pfParticleNet_cff.py @@ -31,7 +31,12 @@ modelConfigPath = cms.FileInPath("RecoBTag/Combined/data/models/particlenet/config.pbtxt"), modelVersion = cms.string(""), verbose = cms.untracked.bool(False), - allowedTries = cms.untracked.uint32(0), + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(0) + ) + ), useSharedMemory = cms.untracked.bool(True), compression = cms.untracked.string(""), ), @@ -64,7 +69,12 @@ modelConfigPath = cms.FileInPath("RecoBTag/Combined/data/models/particlenet_AK8_MD-2prong/config.pbtxt"), modelVersion = cms.string(""), verbose = cms.untracked.bool(False), - allowedTries = cms.untracked.uint32(0), + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(0) + ) + ), ), flav_names = pfMassDecorrelatedParticleNetJetTags.flav_names, )) @@ -94,7 +104,12 @@ modelConfigPath = cms.FileInPath("RecoBTag/Combined/data/models/particlenet_AK8_MassRegression/config.pbtxt"), modelVersion = cms.string(""), verbose = cms.untracked.bool(False), - allowedTries = cms.untracked.uint32(0), + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(0) + ) + ), ), flav_names = pfParticleNetMassRegressionJetTags.flav_names, )) diff --git a/RecoBTag/ONNXRuntime/python/pfParticleTransformerAK4_cff.py b/RecoBTag/ONNXRuntime/python/pfParticleTransformerAK4_cff.py index 3bd1e7761b10d..2aa43e1b5de45 100644 --- a/RecoBTag/ONNXRuntime/python/pfParticleTransformerAK4_cff.py +++ b/RecoBTag/ONNXRuntime/python/pfParticleTransformerAK4_cff.py @@ -18,7 +18,12 @@ modelConfigPath = cms.FileInPath("RecoBTag/Combined/data/models/particletransformer_AK4/config.pbtxt"), modelVersion = cms.string(""), verbose = cms.untracked.bool(False), - allowedTries = cms.untracked.uint32(0), + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(0) + ) + ), useSharedMemory = cms.untracked.bool(True), compression = cms.untracked.string(""), ), diff --git a/RecoBTag/ONNXRuntime/python/pfUnifiedParticleTransformerAK4_cff.py b/RecoBTag/ONNXRuntime/python/pfUnifiedParticleTransformerAK4_cff.py index 6e78ad47b8746..df2c8a6783aef 100644 --- a/RecoBTag/ONNXRuntime/python/pfUnifiedParticleTransformerAK4_cff.py +++ b/RecoBTag/ONNXRuntime/python/pfUnifiedParticleTransformerAK4_cff.py @@ -19,7 +19,12 @@ modelConfigPath = cms.FileInPath("RecoBTag/Combined/data/models/unifiedparticletransformer_AK4_V01/config.pbtxt"), modelVersion = cms.string(""), verbose = cms.untracked.bool(True), - allowedTries = cms.untracked.uint32(0), + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(0) + ) + ), useSharedMemory = cms.untracked.bool(True), compression = cms.untracked.string(""), ), diff --git a/RecoEcal/EgammaClusterProducers/python/SCEnergyCorrectorDRNProducer_cfi.py b/RecoEcal/EgammaClusterProducers/python/SCEnergyCorrectorDRNProducer_cfi.py index 203f1dd49be78..803aa76c11539 100644 --- a/RecoEcal/EgammaClusterProducers/python/SCEnergyCorrectorDRNProducer_cfi.py +++ b/RecoEcal/EgammaClusterProducers/python/SCEnergyCorrectorDRNProducer_cfi.py @@ -6,7 +6,12 @@ mode = cms.string("Async"), modelName = cms.string("MustacheEB"), modelConfigPath = cms.FileInPath("RecoEcal/EgammaClusterProducers/data/models/MustacheEB/config.pbtxt"), - allowedTries = cms.untracked.uint32(1), + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(1) + ) + ), timeout = cms.untracked.uint32(10), ), ) @@ -18,7 +23,12 @@ mode = cms.string("Async"), modelName = cms.string('MustacheEE'), modelConfigPath = cms.FileInPath("RecoEcal/EgammaClusterProducers/data/models/MustacheEE/config.pbtxt"), - allowedTries = cms.untracked.uint32(1), + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(1) + ) + ), timeout = cms.untracked.uint32(10), ), ) diff --git a/RecoTauTag/RecoTau/python/tools/runTauIdMVA.py b/RecoTauTag/RecoTau/python/tools/runTauIdMVA.py index 92a124d26c4a5..6c4d46618bd66 100644 --- a/RecoTauTag/RecoTau/python/tools/runTauIdMVA.py +++ b/RecoTauTag/RecoTau/python/tools/runTauIdMVA.py @@ -797,7 +797,12 @@ def runTauID(self): deepTauSonicTriton.toReplaceWith(_deepTauProducer, DeepTauIdSonicProducer.clone( Client = cms.PSet( mode = cms.string('PseudoAsync'), - allowedTries = cms.untracked.uint32(0), + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(0) + ) + ), verbose = cms.untracked.bool(False), modelName = cms.string("deeptau_2017v2p1"), modelVersion = cms.string(''), @@ -849,7 +854,12 @@ def runTauID(self): deepTauSonicTriton.toReplaceWith(_deepTauProducer, DeepTauIdSonicProducer.clone( Client = cms.PSet( mode = cms.string('PseudoAsync'), - allowedTries = cms.untracked.uint32(0), + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(0) + ) + ), verbose = cms.untracked.bool(False), modelName = cms.string("deeptau_2017v2p1"), modelVersion = cms.string(''), @@ -903,7 +913,12 @@ def runTauID(self): deepTauSonicTriton.toReplaceWith(_deepTauProducer, DeepTauIdSonicProducer.clone( Client = cms.PSet( mode = cms.string('PseudoAsync'), - allowedTries = cms.untracked.uint32(0), + Retry = cms.VPSet( + cms.PSet( + retryType = cms.string('RetrySameServerAction'), + allowedTries = cms.untracked.uint32(0) + ) + ), verbose = cms.untracked.bool(False), modelName = cms.string("deeptau_2018v2p5"), modelVersion = cms.string(''), From c9b55a17050781375f717ebb3f8cfc08e0191509 Mon Sep 17 00:00:00 2001 From: Martin Date: Fri, 24 Jul 2026 16:00:31 -0500 Subject: [PATCH 6/9] Add RetryActionDiff test script. Model needs to be ordered. --- .../SonicTriton/interface/TritonService.h | 2 +- .../SonicTriton/test/BuildFile.xml | 1 + .../SonicTriton/test/retry_action_diff.sh | 39 +++++++++++++++++++ .../SonicTriton/test/retry_action_same.sh | 7 ++-- 4 files changed, 45 insertions(+), 4 deletions(-) create mode 100755 HeterogeneousCore/SonicTriton/test/retry_action_diff.sh diff --git a/HeterogeneousCore/SonicTriton/interface/TritonService.h b/HeterogeneousCore/SonicTriton/interface/TritonService.h index 3be31cb9f2ecd..cf679177e7dd9 100644 --- a/HeterogeneousCore/SonicTriton/interface/TritonService.h +++ b/HeterogeneousCore/SonicTriton/interface/TritonService.h @@ -107,7 +107,7 @@ class TritonService { Model(const std::string& path_ = "") : path(path_) {} //members std::string path; - std::unordered_set servers; + std::set servers; std::unordered_set modules; int refCount{0}; // for dynamic loading on fallback server bool isLoaded() const { return refCount > 0; } diff --git a/HeterogeneousCore/SonicTriton/test/BuildFile.xml b/HeterogeneousCore/SonicTriton/test/BuildFile.xml index 7d5ba1e0acefc..18109d43253c1 100644 --- a/HeterogeneousCore/SonicTriton/test/BuildFile.xml +++ b/HeterogeneousCore/SonicTriton/test/BuildFile.xml @@ -19,6 +19,7 @@ + diff --git a/HeterogeneousCore/SonicTriton/test/retry_action_diff.sh b/HeterogeneousCore/SonicTriton/test/retry_action_diff.sh new file mode 100755 index 0000000000000..59da29ade2d56 --- /dev/null +++ b/HeterogeneousCore/SonicTriton/test/retry_action_diff.sh @@ -0,0 +1,39 @@ +#!/bin/bash + +# Start the server +cmsTriton -v -P 8010 -n server1 -f -L -m /cvmfs/cms.cern.ch/el9_amd64_gcc12/cms/cmssw/CMSSW_15_1_0_pre6/external/el9_amd64_gcc12/data/HeterogeneousCore/SonicTriton/data/models/gat_test/config.pbtxt start & +cmsTriton -v -P 8020 -n server2 -f -L -m /cvmfs/cms.cern.ch/el9_amd64_gcc12/cms/cmssw/CMSSW_15_1_0_pre6/external/el9_amd64_gcc12/data/HeterogeneousCore/SonicTriton/data/models/gat_test/config.pbtxt start & + +# Sleep to allow the server to initialize +sleep 60 + +# Get the true PID of the server +SERVER1_PID=$(ps -e -o pid,cmd | grep "tritonserver" | grep "8010" | grep -v "grep" | awk '{print $1}') +SERVER2_PID=$(ps -e -o pid,cmd | grep "tritonserver" | grep "8020" | grep -v "grep" | awk '{print $1}') + +echo "Servers started with $SERVER1_PID and $SERVER2_PID" +ps aux | grep tritonserver + +if [ -z "$SERVER1_PID" ] || [ -z "$SERVER2_PID" ]; then + echo "Server process could not be started." + exit 1 +fi + +# Start the client +cmsRun tritonTest_cfg.py --maxEvents 100 --modules TritonIdentityProducer --models ragged_io --address server1 0.0.0.0 8011 --address server2 0.0.0.0 8021 --retryAction diff --verbose & + +# Allow the client some time to run +sleep 60 + +# Kill the first server +echo "Killing server with PID $SERVER1_PID" +kill $SERVER1_PID + +# Allow the client some time to run +sleep 30 + +echo "Killing server with PID $SERVER2_PID" +kill $SERVER2_PID +# Wait for client process to complete +wait +echo "Client process completed." diff --git a/HeterogeneousCore/SonicTriton/test/retry_action_same.sh b/HeterogeneousCore/SonicTriton/test/retry_action_same.sh index 1a2c51ba4c4ae..95c14ed6af448 100755 --- a/HeterogeneousCore/SonicTriton/test/retry_action_same.sh +++ b/HeterogeneousCore/SonicTriton/test/retry_action_same.sh @@ -11,14 +11,15 @@ sleep 60 # Get the true PID of the server SERVER_PID=$(ps -e -o pid,cmd | grep "tritonserver" | grep -v "grep" | awk '{print $1}') +echo "Server started with PID $SERVER_PID" + +ps aux | grep tritonserver + if [ -z "$SERVER_PID" ]; then echo "Server process could not be started." exit 1 fi -echo "Server started with PID $SERVER_PID" - -ps aux | grep tritonserver # Start the client cmsRun ${LOCALTOP}/src/HeterogeneousCore/SonicTriton/test/tritonTest_cfg.py --maxEvents 100 --modules TritonIdentityProducer --models ragged_io --address server1 0.0.0.0 8011 --tries 10 --verbose & From e1f2c2fa3d29b338fa3056dfb998587285cbc0e5 Mon Sep 17 00:00:00 2001 From: Martin Date: Fri, 24 Jul 2026 16:19:42 -0500 Subject: [PATCH 7/9] remove commented code --- HeterogeneousCore/SonicTriton/test/retry_action_same.sh | 1 - 1 file changed, 1 deletion(-) diff --git a/HeterogeneousCore/SonicTriton/test/retry_action_same.sh b/HeterogeneousCore/SonicTriton/test/retry_action_same.sh index 95c14ed6af448..309eb3d84faf5 100755 --- a/HeterogeneousCore/SonicTriton/test/retry_action_same.sh +++ b/HeterogeneousCore/SonicTriton/test/retry_action_same.sh @@ -30,7 +30,6 @@ sleep 30 # Kill the server echo "Killing server with PID $SERVER_PID" kill $SERVER_PID -#echo "Server process with PID $SERVER_PID has been killed." # Wait for client process to complete wait From a6e544d6992b0c7b66fe0bb1f6bd4b4e28726fac Mon Sep 17 00:00:00 2001 From: Martin Date: Mon, 3 Aug 2026 14:12:09 -0500 Subject: [PATCH 8/9] update default for the mode parameter --- HeterogeneousCore/SonicCore/README.md | 11 ++++-- .../SonicCore/interface/SonicClientBase.h | 5 ++- .../SonicCore/src/SonicClientBase.cc | 36 ++++++++----------- .../SonicTriton/python/customize.py | 3 -- .../SonicTriton/src/TritonClient.cc | 26 ++++++++++++-- 5 files changed, 51 insertions(+), 30 deletions(-) diff --git a/HeterogeneousCore/SonicCore/README.md b/HeterogeneousCore/SonicCore/README.md index eafcf3ca098ae..a0b0bc5022104 100644 --- a/HeterogeneousCore/SonicCore/README.md +++ b/HeterogeneousCore/SonicCore/README.md @@ -42,7 +42,7 @@ The python configuration for the producer should include a dedicated `PSet` for process.MyProducer = cms.EDProducer("MyProducer", Client = cms.PSet( # necessary client options go here - mode = cms.string("Sync"), + mode = cms.string(""), Retry = cms.VPSet( cms.PSet( retryType = cms.string('RetrySameServerAction'), @@ -104,6 +104,7 @@ The `SonicClient` has three available modes: * `PseudoAsync`: turns a synchronous, blocking call into an asynchronous, non-blocking call, by waiting for the result in a separate `std::thread`. `Async` is the most efficient, but can only be used if asynchronous, non-blocking calls are supported by the communication protocol in use. +When a fallback CPU server is used, `Sync` mode is enforced to avoid contention, user's configuration of `mode` will be respected for other cases with `Async` mode as the default. In addition, as indicated, the input and output data types must be specified. (If both types are the same, only the input type needs to be specified.) @@ -131,6 +132,12 @@ void MyClient::fillPSetDescription(edm::ParameterSetDescription& iDesc) { As indicated, the `fillBasePSetDescription()` function should always be applied to the `descClient` object, to ensure that it includes the necessary parameters. -(Calling `fillBasePSetDescription(descClient, false)` will omit the `Retry` parameter, disabling retries.) +`Retry` is a vector of `PSet` with the parameters `retryType` and `allowedTries`. +`retryType` is the action type that inherents from `RetryActionBase`. +`allowedTries` is a parameter consumed by `RetrySameServerAction`. + +The default `Retry` action is a single `RetryFallbackServerAction`. +To disable `Retry`, leave the `Retry` `VPSet` empty. + Example client code can be found in the `interface` and `src` directories of the other Sonic packages in this repository. diff --git a/HeterogeneousCore/SonicCore/interface/SonicClientBase.h b/HeterogeneousCore/SonicCore/interface/SonicClientBase.h index 45a089701ed12..ed75fa454f2e7 100644 --- a/HeterogeneousCore/SonicCore/interface/SonicClientBase.h +++ b/HeterogeneousCore/SonicCore/interface/SonicClientBase.h @@ -40,10 +40,11 @@ class SonicClientBase { virtual void reset() {} //provide base params - static void fillBasePSetDescription(edm::ParameterSetDescription& desc, bool allowRetry = true); + static void fillBasePSetDescription(edm::ParameterSetDescription& desc); protected: void setMode(SonicMode mode); + void setUserMode(const std::string& userMode); virtual void evaluate() = 0; @@ -70,6 +71,8 @@ class SonicClientBase { //for logging/debugging std::string debugName_, clientName_, fullDebugName_; + //remember what user set at config time + std::string userMode_; friend class SonicDispatcher; friend class SonicDispatcherPseudoAsync; diff --git a/HeterogeneousCore/SonicCore/src/SonicClientBase.cc b/HeterogeneousCore/SonicCore/src/SonicClientBase.cc index 05cee50a2b5f0..4a33ab3605871 100644 --- a/HeterogeneousCore/SonicCore/src/SonicClientBase.cc +++ b/HeterogeneousCore/SonicCore/src/SonicClientBase.cc @@ -9,12 +9,14 @@ void SonicClientBase::RetryDeleter::operator()(RetryActionBase* ptr) const { del SonicClientBase::SonicClientBase(const edm::ParameterSet& params, const std::string& debugName, const std::string& clientName) - : debugName_(debugName), clientName_(clientName), fullDebugName_(debugName_) { + : debugName_(debugName), + clientName_(clientName), + fullDebugName_(debugName_), + userMode_(params.getParameter("mode")) { if (!clientName_.empty()) fullDebugName_ += ":" + clientName_; const auto& retryPSetList = params.getParameter>("Retry"); - std::string modeName(params.getParameter("mode")); for (const auto& retryPSet : retryPSetList) { const std::string& actionType = retryPSet.getParameter("retryType"); @@ -29,16 +31,20 @@ SonicClientBase::SonicClientBase(const edm::ParameterSet& params, } } - if (modeName == "Sync") + setUserMode(userMode_); +} +void SonicClientBase::setUserMode(const std::string& userMode) { + if (userMode == "Sync") setMode(SonicMode::Sync); - else if (modeName == "Async") + else if (userMode == "Async") setMode(SonicMode::Async); - else if (modeName == "PseudoAsync") + else if (userMode == "PseudoAsync") + setMode(SonicMode::PseudoAsync); + else if (userMode == "") setMode(SonicMode::PseudoAsync); else - throw cms::Exception("Configuration") << "Unknown mode for SonicClient: " << modeName; + throw cms::Exception("Configuration") << "Unknown mode for SonicClient: " << userMode; } - void SonicClientBase::setMode(SonicMode mode) { if (dispatcher_ and mode_ == mode) return; @@ -95,24 +101,10 @@ void SonicClientBase::finish(bool success, std::exception_ptr eptr) { reset(); } -void SonicClientBase::fillBasePSetDescription(edm::ParameterSetDescription& desc, bool allowRetry) { +void SonicClientBase::fillBasePSetDescription(edm::ParameterSetDescription& desc) { //restrict allowed values desc.ifValue(edm::ParameterDescription("mode", "PseudoAsync", true), edm::allowedValues("Sync", "Async", "PseudoAsync")); - if (allowRetry) { - // Defines the structure of each entry in the VPSet - edm::ParameterSetDescription retryDesc; - retryDesc.add("retryType", "RetrySameServerAction"); - retryDesc.addUntracked("allowedTries", 0); - - // Define a default retry action - edm::ParameterSet defaultRetry; - defaultRetry.addParameter("retryType", "RetrySameServerAction"); - defaultRetry.addUntrackedParameter("allowedTries", 0); - - // Add the VPSet with the default retry action - desc.addVPSet("Retry", retryDesc, {defaultRetry}); - } desc.add("sonicClientBase", desc); desc.addUntracked("verbose", false); } diff --git a/HeterogeneousCore/SonicTriton/python/customize.py b/HeterogeneousCore/SonicTriton/python/customize.py index e7da1868cd2e8..667d820c22269 100644 --- a/HeterogeneousCore/SonicTriton/python/customize.py +++ b/HeterogeneousCore/SonicTriton/python/customize.py @@ -41,9 +41,6 @@ def getParser(): return parser -def getOptions(parser, verbose=False): - options = parser.parse_args() - def getOptions(parser, verbose=False): options = parser.parse_args() diff --git a/HeterogeneousCore/SonicTriton/src/TritonClient.cc b/HeterogeneousCore/SonicTriton/src/TritonClient.cc index a4781d148c16b..3895073c8efa9 100644 --- a/HeterogeneousCore/SonicTriton/src/TritonClient.cc +++ b/HeterogeneousCore/SonicTriton/src/TritonClient.cc @@ -629,11 +629,19 @@ void TritonClient::updateServer(const std::string& serverName) { serverType_ = server.type; edm::LogInfo("TritonDiscovery") << debugName_ << " assigned server: " << server.url; + //enforce sync mode for fallback CPU server to avoid contention if (serverType_ == TritonServerType::LocalCPU) setMode(SonicMode::Sync); - if (serverType_ == TritonServerType::Remote) - setMode(SonicMode::Async); + else { + if (userMode_.empty()) + // No config from user, default to async for any server type other than localCPU + setMode(SonicMode::Async); + else { + //User configured specific mode for non localCPU server + setUserMode(userMode_); + } + } isLocal_ = serverType_ == TritonServerType::LocalCPU or serverType_ == TritonServerType::LocalGPU; // updateServer() is always called from a TBB thread (via finish() -> retry() -> updateServer()), @@ -680,6 +688,20 @@ void TritonClient::fillPSetDescription(edm::ParameterSetDescription& iDesc) { descClient.addUntracked("useSharedMemory", true); descClient.addUntracked("compression", ""); descClient.addUntracked>("outputs", {}); + + // Defines the structure of each entry in the VPSet (must have a retryType) + edm::ParameterSetDescription retryDesc; + retryDesc.add("retryType", "RetryFallbackServerAction"); + retryDesc.addUntracked("allowedTries", 0); //used by RetrySameServerAction only + + // Define a default retry action + edm::ParameterSet defaultRetry; + defaultRetry.addParameter("retryType", "RetryFallbackServerAction"); + defaultRetry.addUntrackedParameter("allowedTries", 0); + + // Add the VPSet with the default retry action + descClient.addVPSet("Retry", retryDesc, {defaultRetry}); + iDesc.add("Client", descClient); } From 20a104e2acb8a49c8aa727f81c7d93f87ba778a3 Mon Sep 17 00:00:00 2001 From: Martin Date: Wed, 5 Aug 2026 18:15:51 -0500 Subject: [PATCH 9/9] Sonic producers parameters clean up --- .../SonicCore/src/SonicClientBase.cc | 6 +- ..._cfi.py => patElectronDRNCorrector_cff.py} | 9 +-- ...or_cfi.py => patPhotonDRNCorrector_cff.py} | 9 +-- .../PatAlgos/python/slimming/slimming_cff.py | 2 +- .../python/pfHiggsInteractionNet_cff.py | 12 +--- .../python/pfParticleNetAK4_cff.py | 12 +--- .../python/pfParticleNetFromMiniAODAK4_cff.py | 56 +++---------------- .../python/pfParticleNetFromMiniAODAK8_cff.py | 14 +---- .../ONNXRuntime/python/pfParticleNet_cff.py | 38 ++----------- .../python/pfParticleTransformerAK4_cff.py | 14 +---- .../pfUnifiedParticleTransformerAK4_cff.py | 12 +--- .../SCEnergyCorrectorDRNProducer_cff.py | 21 +++++++ .../SCEnergyCorrectorDRNProducer_cfi.py | 36 ------------ .../test/DRNTest_cfg.py | 3 +- .../python/deepMETSonicProducer_cff.py | 1 - .../RecoTau/python/tools/runTauIdMVA.py | 42 +------------- 16 files changed, 52 insertions(+), 235 deletions(-) rename PhysicsTools/PatAlgos/python/slimming/{patElectronDRNCorrector_cfi.py => patElectronDRNCorrector_cff.py} (58%) rename PhysicsTools/PatAlgos/python/slimming/{patPhotonDRNCorrector_cfi.py => patPhotonDRNCorrector_cff.py} (57%) create mode 100644 RecoEcal/EgammaClusterProducers/python/SCEnergyCorrectorDRNProducer_cff.py delete mode 100644 RecoEcal/EgammaClusterProducers/python/SCEnergyCorrectorDRNProducer_cfi.py diff --git a/HeterogeneousCore/SonicCore/src/SonicClientBase.cc b/HeterogeneousCore/SonicCore/src/SonicClientBase.cc index 4a33ab3605871..db67e2a51d751 100644 --- a/HeterogeneousCore/SonicCore/src/SonicClientBase.cc +++ b/HeterogeneousCore/SonicCore/src/SonicClientBase.cc @@ -40,7 +40,7 @@ void SonicClientBase::setUserMode(const std::string& userMode) { setMode(SonicMode::Async); else if (userMode == "PseudoAsync") setMode(SonicMode::PseudoAsync); - else if (userMode == "") + else if (userMode.empty()) setMode(SonicMode::PseudoAsync); else throw cms::Exception("Configuration") << "Unknown mode for SonicClient: " << userMode; @@ -103,8 +103,8 @@ void SonicClientBase::finish(bool success, std::exception_ptr eptr) { void SonicClientBase::fillBasePSetDescription(edm::ParameterSetDescription& desc) { //restrict allowed values - desc.ifValue(edm::ParameterDescription("mode", "PseudoAsync", true), - edm::allowedValues("Sync", "Async", "PseudoAsync")); + desc.ifValue(edm::ParameterDescription("mode", "", true), + edm::allowedValues("Sync", "Async", "PseudoAsync", "")); desc.add("sonicClientBase", desc); desc.addUntracked("verbose", false); } diff --git a/PhysicsTools/PatAlgos/python/slimming/patElectronDRNCorrector_cfi.py b/PhysicsTools/PatAlgos/python/slimming/patElectronDRNCorrector_cff.py similarity index 58% rename from PhysicsTools/PatAlgos/python/slimming/patElectronDRNCorrector_cfi.py rename to PhysicsTools/PatAlgos/python/slimming/patElectronDRNCorrector_cff.py index fa90f55215a6a..8845c047e0e3e 100644 --- a/PhysicsTools/PatAlgos/python/slimming/patElectronDRNCorrector_cfi.py +++ b/PhysicsTools/PatAlgos/python/slimming/patElectronDRNCorrector_cff.py @@ -5,14 +5,7 @@ patElectronsDRN = patElectronDRNCorrectionProducer.clone( particleSource = 'selectedPatElectrons', rhoName = 'fixedGridRhoFastjetAll', - Client = patElectronDRNCorrectionProducer.Client.clone( - mode = 'Async', - Retry = cms.VPSet( - cms.PSet( - retryType = cms.string('RetrySameServerAction'), - allowedTries = cms.untracked.uint32(1) - ) - ), + Client = dict( modelName = 'electronObjectEnsemble', modelConfigPath = 'RecoEgamma/EgammaElectronProducers/data/models/electronObjectEnsemble/config.pbtxt', timeout = 10 diff --git a/PhysicsTools/PatAlgos/python/slimming/patPhotonDRNCorrector_cfi.py b/PhysicsTools/PatAlgos/python/slimming/patPhotonDRNCorrector_cff.py similarity index 57% rename from PhysicsTools/PatAlgos/python/slimming/patPhotonDRNCorrector_cfi.py rename to PhysicsTools/PatAlgos/python/slimming/patPhotonDRNCorrector_cff.py index 4c6c6c2d23230..abc7e1dfd69c9 100644 --- a/PhysicsTools/PatAlgos/python/slimming/patPhotonDRNCorrector_cfi.py +++ b/PhysicsTools/PatAlgos/python/slimming/patPhotonDRNCorrector_cff.py @@ -4,14 +4,7 @@ patPhotonsDRN = patPhotonDRNCorrectionProducer.clone( particleSource = 'selectedPatPhotons', rhoName = 'fixedGridRhoFastjetAll', - Client = patPhotonDRNCorrectionProducer.Client.clone( - mode = 'Async', - Retry = cms.VPSet( - cms.PSet( - retryType = cms.string('RetrySameServerAction'), - allowedTries = cms.untracked.uint32(1) - ) - ), + Client = dict( modelName = 'photonObjectCombined', modelConfigPath = 'RecoEgamma/EgammaPhotonProducers/data/models/photonObjectCombined/config.pbtxt', timeout = 10 diff --git a/PhysicsTools/PatAlgos/python/slimming/slimming_cff.py b/PhysicsTools/PatAlgos/python/slimming/slimming_cff.py index ec7fa573bf6d7..ed8128fa68850 100644 --- a/PhysicsTools/PatAlgos/python/slimming/slimming_cff.py +++ b/PhysicsTools/PatAlgos/python/slimming/slimming_cff.py @@ -117,7 +117,7 @@ offlineSlimmedPrimaryVertices4D) phase2_timing.toReplaceWith(slimmingTask,_phase2_timing_slimmingTask) -from PhysicsTools.PatAlgos.slimming.patPhotonDRNCorrector_cfi import patPhotonsDRN +from PhysicsTools.PatAlgos.slimming.patPhotonDRNCorrector_cff import patPhotonsDRN from Configuration.ProcessModifiers.photonDRN_cff import _photonDRN _photonDRN.toReplaceWith(slimmingTask, cms.Task(slimmingTask.copy(), patPhotonsDRN)) diff --git a/RecoBTag/ONNXRuntime/python/pfHiggsInteractionNet_cff.py b/RecoBTag/ONNXRuntime/python/pfHiggsInteractionNet_cff.py index cc4ef704354cc..5d001e17f9b79 100644 --- a/RecoBTag/ONNXRuntime/python/pfHiggsInteractionNet_cff.py +++ b/RecoBTag/ONNXRuntime/python/pfHiggsInteractionNet_cff.py @@ -26,21 +26,11 @@ particleNetSonicTriton.toReplaceWith(pfHiggsInteractionNetTags, _particleNetSonicJetTagsProducer.clone( src = 'pfHiggsInteractionNetTagInfos', preprocess_json = 'RecoBTag/Combined/data/models/higgsInteractionNet/preprocess.json', - Client = cms.PSet( + Client = dict( timeout = cms.untracked.uint32(300), - mode = cms.string("Async"), modelName = cms.string("higgsInteractionNet"), modelConfigPath = cms.FileInPath("RecoBTag/Combined/data/models/higgsInteractionNet/config.pbtxt"), modelVersion = cms.string(""), - verbose = cms.untracked.bool(False), - Retry = cms.VPSet( - cms.PSet( - retryType = cms.string('RetrySameServerAction'), - allowedTries = cms.untracked.uint32(0) - ) - ), - useSharedMemory = cms.untracked.bool(True), - compression = cms.untracked.string(""), ), flav_names = pfHiggsInteractionNetTags.flav_names, )) diff --git a/RecoBTag/ONNXRuntime/python/pfParticleNetAK4_cff.py b/RecoBTag/ONNXRuntime/python/pfParticleNetAK4_cff.py index 6ec54638152a0..741d91864268b 100644 --- a/RecoBTag/ONNXRuntime/python/pfParticleNetAK4_cff.py +++ b/RecoBTag/ONNXRuntime/python/pfParticleNetAK4_cff.py @@ -41,21 +41,11 @@ particleNetSonicTriton.toReplaceWith(pfParticleNetAK4JetTags, _particleNetSonicJetTagsProducer.clone( src = 'pfParticleNetAK4TagInfos', preprocess_json = 'RecoBTag/Combined/data/ParticleNetAK4/CHS/V00/preprocess.json', - Client = cms.PSet( + Client = dict( timeout = cms.untracked.uint32(300), - mode = cms.string("Async"), modelName = cms.string("particlenet_AK4"), modelConfigPath = cms.FileInPath("RecoBTag/Combined/data/models/particlenet_AK4/config.pbtxt"), modelVersion = cms.string(""), - verbose = cms.untracked.bool(False), - Retry = cms.VPSet( - cms.PSet( - retryType = cms.string('RetrySameServerAction'), - allowedTries = cms.untracked.uint32(0) - ) - ), - useSharedMemory = cms.untracked.bool(True), - compression = cms.untracked.string(""), ), flav_names = pfParticleNetAK4JetTags.flav_names, )) diff --git a/RecoBTag/ONNXRuntime/python/pfParticleNetFromMiniAODAK4_cff.py b/RecoBTag/ONNXRuntime/python/pfParticleNetFromMiniAODAK4_cff.py index ab59555102378..f65d128cc9f57 100644 --- a/RecoBTag/ONNXRuntime/python/pfParticleNetFromMiniAODAK4_cff.py +++ b/RecoBTag/ONNXRuntime/python/pfParticleNetFromMiniAODAK4_cff.py @@ -50,21 +50,11 @@ particleNetSonicTriton.toReplaceWith(pfParticleNetFromMiniAODAK4CHSCentralJetTags, _particleNetSonicJetTagsProducer.clone( src = 'pfParticleNetFromMiniAODAK4CHSCentralTagInfos', preprocess_json = 'RecoBTag/Combined/data/ParticleNetFromMiniAODAK4/CHS/Central/preprocess.json', - Client = cms.PSet( + Client = dict( timeout = cms.untracked.uint32(300), - mode = cms.string("Async"), modelName = cms.string("particleNetFromMiniAODAK4CHSCentral"), modelConfigPath = cms.FileInPath("RecoBTag/Combined/data/models/particleNetFromMiniAODAK4CHSCentral/config.pbtxt"), - modelVersion = cms.string(""), - verbose = cms.untracked.bool(False), - Retry = cms.VPSet( - cms.PSet( - retryType = cms.string('RetrySameServerAction'), - allowedTries = cms.untracked.uint32(0) - ) - ), - useSharedMemory = cms.untracked.bool(True), - compression = cms.untracked.string(""), + modelVersion = cms.string("") ), flav_names = pfParticleNetFromMiniAODAK4CHSCentralJetTags.flav_names, )) @@ -79,21 +69,11 @@ particleNetSonicTriton.toReplaceWith(pfParticleNetFromMiniAODAK4CHSForwardJetTags, _particleNetSonicJetTagsProducer.clone( src = 'pfParticleNetFromMiniAODAK4CHSForwardTagInfos', preprocess_json = 'RecoBTag/Combined/data/ParticleNetFromMiniAODAK4/CHS/Central/preprocess.json', - Client = cms.PSet( + Client = dict( timeout = cms.untracked.uint32(300), - mode = cms.string("Async"), modelName = cms.string("particleNetFromMiniAODAK4CHSForward"), modelConfigPath = cms.FileInPath("RecoBTag/Combined/data/models/particleNetFromMiniAODAK4CHSForward/config.pbtxt"), - modelVersion = cms.string(""), - verbose = cms.untracked.bool(False), - Retry = cms.VPSet( - cms.PSet( - retryType = cms.string('RetrySameServerAction'), - allowedTries = cms.untracked.uint32(0) - ) - ), - useSharedMemory = cms.untracked.bool(True), - compression = cms.untracked.string(""), + modelVersion = cms.string("") ), flav_names = pfParticleNetFromMiniAODAK4CHSForwardJetTags.flav_names, )) @@ -108,21 +88,11 @@ particleNetSonicTriton.toReplaceWith(pfParticleNetFromMiniAODAK4PuppiCentralJetTags, _particleNetSonicJetTagsProducer.clone( src = 'pfParticleNetFromMiniAODAK4PuppiCentralTagInfos', preprocess_json = 'RecoBTag/Combined/data/ParticleNetFromMiniAODAK4/PUPPI/Central/preprocess.json', - Client = cms.PSet( + Client = dict( timeout = cms.untracked.uint32(300), - mode = cms.string("Async"), modelName = cms.string("particleNetFromMiniAODAK4PuppiCentral"), modelConfigPath = cms.FileInPath("RecoBTag/Combined/data/models/particleNetFromMiniAODAK4PuppiCentral/config.pbtxt"), - modelVersion = cms.string(""), - verbose = cms.untracked.bool(False), - Retry = cms.VPSet( - cms.PSet( - retryType = cms.string('RetrySameServerAction'), - allowedTries = cms.untracked.uint32(0) - ) - ), - useSharedMemory = cms.untracked.bool(True), - compression = cms.untracked.string(""), + modelVersion = cms.string("") ), flav_names = pfParticleNetFromMiniAODAK4PuppiCentralJetTags.flav_names, )) @@ -137,21 +107,11 @@ particleNetSonicTriton.toReplaceWith(pfParticleNetFromMiniAODAK4PuppiForwardJetTags, _particleNetSonicJetTagsProducer.clone( src = 'pfParticleNetFromMiniAODAK4PuppiForwardTagInfos', preprocess_json = 'RecoBTag/Combined/data/ParticleNetFromMiniAODAK4/PUPPI/Forward/preprocess.json', - Client = cms.PSet( + Client = dict( timeout = cms.untracked.uint32(300), - mode = cms.string("Async"), modelName = cms.string("particleNetFromMiniAODAK4PuppiForward"), modelConfigPath = cms.FileInPath("RecoBTag/Combined/data/models/particleNetFromMiniAODAK4PuppiForward/config.pbtxt"), - modelVersion = cms.string(""), - verbose = cms.untracked.bool(False), - Retry = cms.VPSet( - cms.PSet( - retryType = cms.string('RetrySameServerAction'), - allowedTries = cms.untracked.uint32(0) - ) - ), - useSharedMemory = cms.untracked.bool(True), - compression = cms.untracked.string(""), + modelVersion = cms.string("") ), flav_names = pfParticleNetFromMiniAODAK4PuppiForwardJetTags.flav_names, )) diff --git a/RecoBTag/ONNXRuntime/python/pfParticleNetFromMiniAODAK8_cff.py b/RecoBTag/ONNXRuntime/python/pfParticleNetFromMiniAODAK8_cff.py index 529dfce82eb3d..7eb661f15ceec 100644 --- a/RecoBTag/ONNXRuntime/python/pfParticleNetFromMiniAODAK8_cff.py +++ b/RecoBTag/ONNXRuntime/python/pfParticleNetFromMiniAODAK8_cff.py @@ -26,21 +26,11 @@ particleNetSonicTriton.toReplaceWith(pfParticleNetFromMiniAODAK8JetTags, _particleNetSonicJetTagsProducer.clone( src = 'pfParticleNetFromMiniAODAK8TagInfos', preprocess_json = 'RecoBTag/Combined/data/ParticleNetFromMiniAODAK8/preprocess.json', - Client = cms.PSet( + Client = dict( timeout = cms.untracked.uint32(300), - mode = cms.string("Async"), modelName = cms.string("particleNetFromMiniAODAK8"), modelConfigPath = cms.FileInPath("RecoBTag/Combined/data/models/particleNetFromMiniAODAK8/config.pbtxt"), - modelVersion = cms.string(""), - verbose = cms.untracked.bool(False), - Retry = cms.VPSet( - cms.PSet( - retryType = cms.string('RetrySameServerAction'), - allowedTries = cms.untracked.uint32(0) - ) - ), - useSharedMemory = cms.untracked.bool(True), - compression = cms.untracked.string(""), + modelVersion = cms.string("") ), flav_names = pfParticleNetFromMiniAODAK8JetTags.flav_names, )) diff --git a/RecoBTag/ONNXRuntime/python/pfParticleNet_cff.py b/RecoBTag/ONNXRuntime/python/pfParticleNet_cff.py index 93df39ca9c9b3..b3b8c2969f54d 100644 --- a/RecoBTag/ONNXRuntime/python/pfParticleNet_cff.py +++ b/RecoBTag/ONNXRuntime/python/pfParticleNet_cff.py @@ -24,21 +24,11 @@ particleNetSonicTriton.toReplaceWith(pfParticleNetJetTags, _particleNetSonicJetTagsProducer.clone( src = 'pfParticleNetTagInfos', preprocess_json = 'RecoBTag/Combined/data/ParticleNetAK8/General/V01/preprocess.json', - Client = cms.PSet( + Client = dict( timeout = cms.untracked.uint32(300), - mode = cms.string("Async"), modelName = cms.string("particlenet"), modelConfigPath = cms.FileInPath("RecoBTag/Combined/data/models/particlenet/config.pbtxt"), - modelVersion = cms.string(""), - verbose = cms.untracked.bool(False), - Retry = cms.VPSet( - cms.PSet( - retryType = cms.string('RetrySameServerAction'), - allowedTries = cms.untracked.uint32(0) - ) - ), - useSharedMemory = cms.untracked.bool(True), - compression = cms.untracked.string(""), + modelVersion = cms.string("") ), flav_names = pfParticleNetJetTags.flav_names, )) @@ -62,19 +52,11 @@ particleNetSonicTriton.toReplaceWith(pfMassDecorrelatedParticleNetJetTags, _particleNetSonicJetTagsProducer.clone( src = 'pfParticleNetTagInfos', preprocess_json = 'RecoBTag/Combined/data/ParticleNetAK8/MD-2prong/V01/preprocess.json', - Client = cms.PSet( + Client = dict( timeout = cms.untracked.uint32(300), modelName = cms.string("particlenet_AK8_MD-2prong"), - mode = cms.string("Async"), modelConfigPath = cms.FileInPath("RecoBTag/Combined/data/models/particlenet_AK8_MD-2prong/config.pbtxt"), - modelVersion = cms.string(""), - verbose = cms.untracked.bool(False), - Retry = cms.VPSet( - cms.PSet( - retryType = cms.string('RetrySameServerAction'), - allowedTries = cms.untracked.uint32(0) - ) - ), + modelVersion = cms.string("") ), flav_names = pfMassDecorrelatedParticleNetJetTags.flav_names, )) @@ -97,19 +79,11 @@ particleNetSonicTriton.toReplaceWith(pfParticleNetMassRegressionJetTags, _particleNetSonicJetTagsProducer.clone( src = 'pfParticleNetTagInfos', preprocess_json = 'RecoBTag/Combined/data/ParticleNetAK8/MassRegression/V01/preprocess.json', - Client = cms.PSet( + Client = dict( timeout = cms.untracked.uint32(300), modelName = cms.string("particlenet_AK8_MassRegression"), - mode = cms.string("Async"), modelConfigPath = cms.FileInPath("RecoBTag/Combined/data/models/particlenet_AK8_MassRegression/config.pbtxt"), - modelVersion = cms.string(""), - verbose = cms.untracked.bool(False), - Retry = cms.VPSet( - cms.PSet( - retryType = cms.string('RetrySameServerAction'), - allowedTries = cms.untracked.uint32(0) - ) - ), + modelVersion = cms.string("") ), flav_names = pfParticleNetMassRegressionJetTags.flav_names, )) diff --git a/RecoBTag/ONNXRuntime/python/pfParticleTransformerAK4_cff.py b/RecoBTag/ONNXRuntime/python/pfParticleTransformerAK4_cff.py index 2aa43e1b5de45..80dcda4335d56 100644 --- a/RecoBTag/ONNXRuntime/python/pfParticleTransformerAK4_cff.py +++ b/RecoBTag/ONNXRuntime/python/pfParticleTransformerAK4_cff.py @@ -11,21 +11,11 @@ particleTransformerAK4SonicTriton.toReplaceWith(pfParticleTransformerAK4JetTags, _pfParticleTransformerAK4SonicJetTags.clone( - Client = cms.PSet( + Client = dict( timeout = cms.untracked.uint32(300), - mode = cms.string("Async"), modelName = cms.string("particletransformer_AK4"), modelConfigPath = cms.FileInPath("RecoBTag/Combined/data/models/particletransformer_AK4/config.pbtxt"), - modelVersion = cms.string(""), - verbose = cms.untracked.bool(False), - Retry = cms.VPSet( - cms.PSet( - retryType = cms.string('RetrySameServerAction'), - allowedTries = cms.untracked.uint32(0) - ) - ), - useSharedMemory = cms.untracked.bool(True), - compression = cms.untracked.string(""), + modelVersion = cms.string("") ), flav_names = pfParticleTransformerAK4JetTags.flav_names, )) diff --git a/RecoBTag/ONNXRuntime/python/pfUnifiedParticleTransformerAK4_cff.py b/RecoBTag/ONNXRuntime/python/pfUnifiedParticleTransformerAK4_cff.py index df2c8a6783aef..e89200acca996 100644 --- a/RecoBTag/ONNXRuntime/python/pfUnifiedParticleTransformerAK4_cff.py +++ b/RecoBTag/ONNXRuntime/python/pfUnifiedParticleTransformerAK4_cff.py @@ -12,21 +12,11 @@ pfUnifiedParticleTransformerAK4JetTags = _pfUnifiedParticleTransformerAK4JetTags.clone() unifiedparticleTransformerAK4SonicTriton.toReplaceWith(pfUnifiedParticleTransformerAK4JetTags, _pfUnifiedParticleTransformerAK4SonicJetTags.clone( - Client = cms.PSet( + Client = dict( timeout = cms.untracked.uint32(500), - mode = cms.string("Async"), modelName = cms.string("unifiedparticletransformer_AK4_V01"), modelConfigPath = cms.FileInPath("RecoBTag/Combined/data/models/unifiedparticletransformer_AK4_V01/config.pbtxt"), modelVersion = cms.string(""), - verbose = cms.untracked.bool(True), - Retry = cms.VPSet( - cms.PSet( - retryType = cms.string('RetrySameServerAction'), - allowedTries = cms.untracked.uint32(0) - ) - ), - useSharedMemory = cms.untracked.bool(True), - compression = cms.untracked.string(""), ), flav_names = pfUnifiedParticleTransformerAK4JetTags.flav_names, )) diff --git a/RecoEcal/EgammaClusterProducers/python/SCEnergyCorrectorDRNProducer_cff.py b/RecoEcal/EgammaClusterProducers/python/SCEnergyCorrectorDRNProducer_cff.py new file mode 100644 index 0000000000000..f330bf41e46d1 --- /dev/null +++ b/RecoEcal/EgammaClusterProducers/python/SCEnergyCorrectorDRNProducer_cff.py @@ -0,0 +1,21 @@ +import FWCore.ParameterSet.Config as cms + +from RecoEcal.EgammaClusterProducers.SCEnergyCorrectorDRNProducer_cfi import SCEnergyCorrectorDRNProducer as _SCEnergyCorrectorDRNProducer + +DRNProducerEB = _SCEnergyCorrectorDRNProducer.clone( + inputSCs = "particleFlowSuperClusterECAL:particleFlowSuperClusterECALBarrel", + Client = dict( + modelName = "MustacheEB", + modelConfigPath = "RecoEcal/EgammaClusterProducers/data/models/MustacheEB/config.pbtxt", + timeout = 10, + ), +) + +DRNProducerEE = _SCEnergyCorrectorDRNProducer.clone( + inputSCs = "particleFlowSuperClusterECAL:particleFlowSuperClusterECALEndcapWithPreshower", + Client = dict( + modelName = "MustacheEE", + modelConfigPath = "RecoEcal/EgammaClusterProducers/data/models/MustacheEE/config.pbtxt", + timeout = 10, + ), +) diff --git a/RecoEcal/EgammaClusterProducers/python/SCEnergyCorrectorDRNProducer_cfi.py b/RecoEcal/EgammaClusterProducers/python/SCEnergyCorrectorDRNProducer_cfi.py deleted file mode 100644 index 803aa76c11539..0000000000000 --- a/RecoEcal/EgammaClusterProducers/python/SCEnergyCorrectorDRNProducer_cfi.py +++ /dev/null @@ -1,36 +0,0 @@ -import FWCore.ParameterSet.Config as cms - -DRNProducerEB = cms.EDProducer('SCEnergyCorrectorDRNProducer', - inputSCs = cms.InputTag('particleFlowSuperClusterECAL','particleFlowSuperClusterECALBarrel'), - Client = cms.PSet( - mode = cms.string("Async"), - modelName = cms.string("MustacheEB"), - modelConfigPath = cms.FileInPath("RecoEcal/EgammaClusterProducers/data/models/MustacheEB/config.pbtxt"), - Retry = cms.VPSet( - cms.PSet( - retryType = cms.string('RetrySameServerAction'), - allowedTries = cms.untracked.uint32(1) - ) - ), - timeout = cms.untracked.uint32(10), - ), - ) - - -DRNProducerEE = cms.EDProducer('SCEnergyCorrectorDRNProducer', - inputSCs = cms.InputTag('particleFlowSuperClusterECAL','particleFlowSuperClusterECALEndcapWithPreshower'), - Client = cms.PSet( - mode = cms.string("Async"), - modelName = cms.string('MustacheEE'), - modelConfigPath = cms.FileInPath("RecoEcal/EgammaClusterProducers/data/models/MustacheEE/config.pbtxt"), - Retry = cms.VPSet( - cms.PSet( - retryType = cms.string('RetrySameServerAction'), - allowedTries = cms.untracked.uint32(1) - ) - ), - timeout = cms.untracked.uint32(10), - ), - ) - - diff --git a/RecoEcal/EgammaClusterProducers/test/DRNTest_cfg.py b/RecoEcal/EgammaClusterProducers/test/DRNTest_cfg.py index f777e379362f7..98025eb073cc3 100644 --- a/RecoEcal/EgammaClusterProducers/test/DRNTest_cfg.py +++ b/RecoEcal/EgammaClusterProducers/test/DRNTest_cfg.py @@ -36,8 +36,7 @@ from Configuration.AlCa.GlobalTag import GlobalTag process.GlobalTag = GlobalTag(process.GlobalTag, '106X_upgrade2018_realistic_v11_Ecal5', '') -#from RecoEcal.EgammaClusterProducers.SCEnergyCorrectorDRNProducer_cfi import * -from RecoEcal.EgammaClusterProducers.SCEnergyCorrectorDRNProducer_cfi import * +from RecoEcal.EgammaClusterProducers.SCEnergyCorrectorDRNProducer_cff import * process.DRNProducerEB = DRNProducerEB process.DRNProducerEE = DRNProducerEE diff --git a/RecoMET/METPUSubtraction/python/deepMETSonicProducer_cff.py b/RecoMET/METPUSubtraction/python/deepMETSonicProducer_cff.py index 218fe08d1d35d..23c7b299fbbcf 100644 --- a/RecoMET/METPUSubtraction/python/deepMETSonicProducer_cff.py +++ b/RecoMET/METPUSubtraction/python/deepMETSonicProducer_cff.py @@ -5,7 +5,6 @@ deepMETSonicProducer = _deepMETSonicProducer.clone( Client = dict( timeout = 300, - mode = "Async", modelName = "deepmet", modelConfigPath = "RecoMET/METPUSubtraction/data/models/deepmet/config.pbtxt", # version "1" is the resolutionTune diff --git a/RecoTauTag/RecoTau/python/tools/runTauIdMVA.py b/RecoTauTag/RecoTau/python/tools/runTauIdMVA.py index 6c4d46618bd66..37440f80decdb 100644 --- a/RecoTauTag/RecoTau/python/tools/runTauIdMVA.py +++ b/RecoTauTag/RecoTau/python/tools/runTauIdMVA.py @@ -795,23 +795,11 @@ def runTauID(self): from Configuration.ProcessModifiers.deepTauSonicTriton_cff import deepTauSonicTriton deepTauSonicTriton.toReplaceWith(_deepTauProducer, DeepTauIdSonicProducer.clone( - Client = cms.PSet( - mode = cms.string('PseudoAsync'), - Retry = cms.VPSet( - cms.PSet( - retryType = cms.string('RetrySameServerAction'), - allowedTries = cms.untracked.uint32(0) - ) - ), - verbose = cms.untracked.bool(False), + Client = dict( modelName = cms.string("deeptau_2017v2p1"), modelVersion = cms.string(''), modelConfigPath = cms.FileInPath("RecoTauTag/TrainingFiles/data/DeepTauIdSONIC/deeptau_2017v2p1/config.pbtxt"), - preferredServer = cms.untracked.string(''), timeout = cms.untracked.uint32(300), - useSharedMemory = cms.untracked.bool(True), - compression = cms.untracked.string(''), - outputs = cms.untracked.vstring() ), Prediscriminants = noPrediscriminants, taus = self.originalTauName, @@ -852,23 +840,11 @@ def runTauID(self): from Configuration.ProcessModifiers.deepTauSonicTriton_cff import deepTauSonicTriton deepTauSonicTriton.toReplaceWith(_deepTauProducer, DeepTauIdSonicProducer.clone( - Client = cms.PSet( - mode = cms.string('PseudoAsync'), - Retry = cms.VPSet( - cms.PSet( - retryType = cms.string('RetrySameServerAction'), - allowedTries = cms.untracked.uint32(0) - ) - ), - verbose = cms.untracked.bool(False), + Client = dict( modelName = cms.string("deeptau_2017v2p1"), modelVersion = cms.string(''), modelConfigPath = cms.FileInPath("RecoTauTag/TrainingFiles/data/DeepTauIdSONIC/deeptau_2017v2p1/config.pbtxt"), - preferredServer = cms.untracked.string(''), timeout = cms.untracked.uint32(300), - useSharedMemory = cms.untracked.bool(True), - compression = cms.untracked.string(''), - outputs = cms.untracked.vstring() ), Prediscriminants = noPrediscriminants, taus = self.originalTauName, @@ -911,23 +887,11 @@ def runTauID(self): from Configuration.ProcessModifiers.deepTauSonicTriton_cff import deepTauSonicTriton deepTauSonicTriton.toReplaceWith(_deepTauProducer, DeepTauIdSonicProducer.clone( - Client = cms.PSet( - mode = cms.string('PseudoAsync'), - Retry = cms.VPSet( - cms.PSet( - retryType = cms.string('RetrySameServerAction'), - allowedTries = cms.untracked.uint32(0) - ) - ), - verbose = cms.untracked.bool(False), + Client = dict( modelName = cms.string("deeptau_2018v2p5"), modelVersion = cms.string(''), modelConfigPath = cms.FileInPath("RecoTauTag/TrainingFiles/data/DeepTauIdSONIC/deeptau_2018v2p5/config.pbtxt"), - preferredServer = cms.untracked.string(''), timeout = cms.untracked.uint32(300), - useSharedMemory = cms.untracked.bool(True), - compression = cms.untracked.string(''), - outputs = cms.untracked.vstring(), ), Prediscriminants = noPrediscriminants, taus = self.originalTauName,