ATLAS Offline Software
Loading...
Searching...
No Matches
AthInfer::TritonTool::Impl Struct Reference
Inheritance diagram for AthInfer::TritonTool::Impl:
Collaboration diagram for AthInfer::TritonTool::Impl:

Public Member Functions

StatusCode getClient (tc::InferenceServerGrpcClient *&client, const std::string &url, int port, bool useSSL) const
tc::Error checkServerHealth (tc::InferenceServerGrpcClient &client) const
void waitBeforeRetry (int retryDelayMs, int attempt) const
StatusCode runInference (tc::InferenceServerGrpcClient &client, const std::vector< tc::InferInput * > &rawInputs, const int maxRetries, const int retryDelayMs, std::shared_ptr< tc::InferResult > &results) const
template<typename T>
StatusCode prepareInput (const std::string &name, const std::vector< int64_t > &shape, const std::vector< T > &data, std::vector< std::unique_ptr< tc::InferInput > > &inputs) const
template<typename T>
StatusCode extractOutput (const std::string &name, const tc::InferResult &result, std::vector< T > &outputVec) const
 AthMessaging (IMessageSvc *msgSvc, const std::string &name)
 Constructor.
 AthMessaging (const std::string &name)
 Constructor with auto-retrieval of the MessageSvc.
bool msgLvl (const MSG::Level lvl) const
 Test the output level.
MsgStream & msg () const
 The standard message stream.
MsgStream & msg (const MSG::Level lvl) const
 The standard message stream.
void setLevel (MSG::Level lvl)
 Change the current logging level.

Public Attributes

const AthAsynchronousAlgorithmm_parentAsyncAlg = nullptr
std::unique_ptr< tc::InferOptions > m_options

Private Member Functions

void initMessaging () const
 Initialize our message level and MessageSvc.

Private Attributes

std::string m_nm
 Message source name.
boost::thread_specific_ptr< MsgStream > m_msg_tls
 MsgStream instance (a std::cout like with print-out levels).
std::atomic< IMessageSvc * > m_imsg { nullptr }
 MessageSvc pointer.
std::atomic< MSG::Level > m_lvl { MSG::NIL }
 Current logging level.
std::atomic_flag m_initialized ATLAS_THREAD_SAFE = ATOMIC_FLAG_INIT
 Messaging initialized (initMessaging).

Detailed Description

Definition at line 55 of file TritonTool.cxx.

Member Function Documentation

◆ AthMessaging() [1/2]

AthMessaging::AthMessaging ( const std::string & name)

Constructor with auto-retrieval of the MessageSvc.

Parameters
nameName of the message stream

Definition at line 71 of file AthMessaging.cxx.

19 :
20 m_nm(name)
21{}
std::string m_nm
Message source name.

◆ AthMessaging() [2/2]

AthMessaging::AthMessaging ( IMessageSvc * msgSvc,
const std::string & name )

Constructor.

Parameters
msgSvcPointer to the MessageSvc
nameName of the message stream

Definition at line 66 of file AthMessaging.cxx.

14 :
15 m_nm(name), m_imsg(msgSvc)
16{}
std::atomic< IMessageSvc * > m_imsg
MessageSvc pointer.

◆ checkServerHealth()

tc::Error AthInfer::TritonTool::Impl::checkServerHealth ( tc::InferenceServerGrpcClient & client) const
inline

Definition at line 80 of file TritonTool.cxx.

80 {
81
82 tc::Headers httpHeaders;
83 bool live = false;
84 tc::Error err =
85 client.IsServerLive(&live, httpHeaders, m_options->client_timeout_);
86 if (!err.IsOk()) {
87 return err;
88 }
89 if (!live) {
90 return tc::Error("Triton server is not live");
91 }
92
93 bool serverReady = false;
94 err = client.IsServerReady(&serverReady, httpHeaders,
95 m_options->client_timeout_);
96 if (!err.IsOk()) {
97 return err;
98 }
99 if (!serverReady) {
100 return tc::Error("Triton server is not ready");
101 }
102
103 bool modelReady = false;
104 err = client.IsModelReady(&modelReady, m_options->model_name_,
105 m_options->model_version_, httpHeaders,
106 m_options->client_timeout_);
107 if (!err.IsOk()) {
108 return err;
109 }
110 if (!modelReady) {
111 return tc::Error("Triton model " + m_options->model_name_ + " is not ready");
112 }
113
114 return tc::Error::Success;
115 }
std::unique_ptr< tc::InferOptions > m_options

◆ extractOutput()

template<typename T>
StatusCode AthInfer::TritonTool::Impl::extractOutput ( const std::string & name,
const tc::InferResult & result,
std::vector< T > & outputVec ) const
inline

Definition at line 217 of file TritonTool.cxx.

219 {
220
221 const uint8_t* rawData = nullptr;
222 size_t size = 0;
223
224 // Get access to the buffer holding raw results of specified output returned
225 // by the server. Note: the buffer is owned by InferResult instance. Users
226 // can copy out the data if required to extend the lifetime.
227 TRITON_CHECK(result.RawData(name, &rawData, &size));
228
229 outputVec.resize(size / sizeof(T));
230 std::memcpy(outputVec.data(), rawData, size);
231 return StatusCode::SUCCESS;
232 }
size_t size() const
Number of registered mappings.
#define TRITON_CHECK(EXP)
Shorthand for the Triton client namespace.

◆ getClient()

StatusCode AthInfer::TritonTool::Impl::getClient ( tc::InferenceServerGrpcClient *& client,
const std::string & url,
int port,
bool useSSL ) const
inline

Definition at line 60 of file TritonTool.cxx.

61 {
62
63 thread_local std::unique_ptr<tc::InferenceServerGrpcClient> threadClient;
64 if (!threadClient) {
65
66 const std::string urlAndPort =
67 url + ":" + std::to_string(port); // always use the gRPC port
68
69 constexpr bool verbose = false;
70 TRITON_CHECK(tc::InferenceServerGrpcClient::Create(
71 &threadClient, urlAndPort, verbose, useSSL));
72
73 ATH_MSG_INFO("Triton client created for url: " << urlAndPort);
74 }
75 client = threadClient.get();
76
77 return StatusCode::SUCCESS;
78 }
#define ATH_MSG_INFO(x,...)
bool verbose
Definition hcg.cxx:75

◆ initMessaging()

void AthMessaging::initMessaging ( ) const
privateinherited

Initialize our message level and MessageSvc.

This method should only be called once.

Definition at line 39 of file AthMessaging.cxx.

40{
42 // If user did not set an explicit level, set a default
43 if (m_lvl == MSG::NIL) {
44 m_lvl = m_imsg ?
45 static_cast<MSG::Level>( m_imsg.load()->outputLevel(m_nm) ) :
46 MSG::INFO;
47 }
48}
std::atomic< MSG::Level > m_lvl
Current logging level.
IMessageSvc * getMessageSvc(bool quiet=false)
int outputLevel(const IMessageSvc *ims, const std::string &source)

◆ msg() [1/2]

MsgStream & AthMessaging::msg ( ) const
inlineinherited

The standard message stream.

Returns a reference to the default message stream May not be invoked before sysInitialize() has been invoked.

Definition at line 167 of file AthMessaging.h.

168{
169 MsgStream* ms = m_msg_tls.get();
170 if (!ms) {
171 if (!m_initialized.test_and_set()) initMessaging();
172 ms = new MsgStream(m_imsg,m_nm);
173 m_msg_tls.reset( ms );
174 }
175
176 ms->setLevel (m_lvl);
177 return *ms;
178}
boost::thread_specific_ptr< MsgStream > m_msg_tls
MsgStream instance (a std::cout like with print-out levels).
void initMessaging() const
Initialize our message level and MessageSvc.

◆ msg() [2/2]

MsgStream & AthMessaging::msg ( const MSG::Level lvl) const
inlineinherited

The standard message stream.

Returns a reference to the default message stream May not be invoked before sysInitialize() has been invoked.

Definition at line 182 of file AthMessaging.h.

183{ return msg() << lvl; }
MsgStream & msg() const
The standard message stream.

◆ msgLvl()

bool AthMessaging::msgLvl ( const MSG::Level lvl) const
inlineinherited

Test the output level.

Parameters
lvlThe message level to test against
Returns
boolean Indicating if messages at given level will be printed
Return values
trueMessages at level "lvl" will be printed

Definition at line 151 of file AthMessaging.h.

152{
153 // If user did not set explicit message level we have to initialize
154 // the messaging and retrieve the default via the MessageSvc.
155 if (m_lvl==MSG::NIL && !m_initialized.test_and_set()) initMessaging();
156
157 if (m_lvl <= lvl) {
158 msg() << lvl;
159 return true;
160 } else {
161 return false;
162 }
163}

◆ prepareInput()

template<typename T>
StatusCode AthInfer::TritonTool::Impl::prepareInput ( const std::string & name,
const std::vector< int64_t > & shape,
const std::vector< T > & data,
std::vector< std::unique_ptr< tc::InferInput > > & inputs ) const
inline

Definition at line 189 of file TritonTool.cxx.

192 {
193
194 const char* dtype = TritonDType<T>::value;
195 tc::InferInput* rawInputPtr = nullptr;
196
197 // create the InferInput object with the predefined name, shape, and data
198 // type.
199 TRITON_CHECK(tc::InferInput::Create(&rawInputPtr, name, shape, dtype));
200 assert(rawInputPtr != nullptr);
201
202 // Append tensor values for this input from a byte array.
203 // Note: The vector is not copied and so it must not be modified or
204 // destroyed until this input is no longer needed (that is until the Infer()
205 // call(s) that use the input have completed). Multiple calls can be made to
206 // this API to keep adding tensor data for this input. The data will be
207 // delivered in the order it was added.
208 std::unique_ptr<tc::InferInput> input{rawInputPtr};
209 TRITON_CHECK(input->AppendRaw(reinterpret_cast<const uint8_t*>(data.data()),
210 data.size() * sizeof(T)));
211
212 inputs.push_back(std::move(input));
213 return StatusCode::SUCCESS;
214 }

◆ runInference()

StatusCode AthInfer::TritonTool::Impl::runInference ( tc::InferenceServerGrpcClient & client,
const std::vector< tc::InferInput * > & rawInputs,
const int maxRetries,
const int retryDelayMs,
std::shared_ptr< tc::InferResult > & results ) const
inline

Definition at line 124 of file TritonTool.cxx.

128 {
129
130 tc::Headers httpHeaders;
131 grpc_compression_algorithm compressionAlgorithm =
132 grpc_compression_algorithm::GRPC_COMPRESS_NONE;
133
134 tc::Error err;
135 for (int attempt = 0; attempt <= maxRetries; ++attempt) {
136 if (m_parentAsyncAlg == nullptr) {
137 tc::InferResult* rawResultPtr = nullptr;
138 err = client.Infer(&rawResultPtr, *m_options, rawInputs, {},
139 httpHeaders, compressionAlgorithm);
140 if (err.IsOk() && rawResultPtr != nullptr) {
141 results.reset(rawResultPtr);
142 err = results->RequestStatus();
143 } else if (err.IsOk()) {
144 err = tc::Error("Triton synchronous inference returned no result");
145 }
146 } else {
147 using Promise_t = boost::fibers::promise<tc::InferResult*>;
148 using Future_t = boost::fibers::future<tc::InferResult*>;
149 Promise_t promise{};
150 Future_t future = promise.get_future();
151 auto callback = [&promise](tc::InferResult* resultPtr) {
152 promise.set_value(resultPtr);
153 };
154 err = client.AsyncInfer(callback, *m_options, rawInputs, {},
155 httpHeaders, compressionAlgorithm);
156 if (err.IsOk()) {
157 results.reset(future.get());
158 ATH_CHECK(m_parentAsyncAlg->restoreAfterSuspend());
159 if (results != nullptr) {
160 err = results->RequestStatus();
161 } else {
162 err = tc::Error("Triton asynchronous inference returned no "
163 "result");
164 }
165 }
166 }
167
168 if (err.IsOk()) {
169 return StatusCode::SUCCESS;
170 }
171
172 if (attempt == maxRetries) {
173 ATH_MSG_ERROR("Triton inference failed after " << (attempt + 1)
174 << " attempt(s): "
175 << err);
176 return StatusCode::FAILURE;
177 }
178
179 ATH_MSG_WARNING("Triton inference attempt " << (attempt + 1)
180 << " failed: " << err
181 << "; retrying");
182 waitBeforeRetry(retryDelayMs, attempt);
183 }
184
185 return StatusCode::FAILURE;
186 }
#define ATH_CHECK
Evaluate an expression and check for errors.
#define ATH_MSG_ERROR(x,...)
#define ATH_MSG_WARNING(x,...)
const AthAsynchronousAlgorithm * m_parentAsyncAlg
void waitBeforeRetry(int retryDelayMs, int attempt) const

◆ setLevel()

void AthMessaging::setLevel ( MSG::Level lvl)
inherited

Change the current logging level.

Use this rather than msg().setLevel() for proper operation with MT.

Definition at line 28 of file AthMessaging.cxx.

29{
30 m_lvl = lvl;
31}

◆ waitBeforeRetry()

void AthInfer::TritonTool::Impl::waitBeforeRetry ( int retryDelayMs,
int attempt ) const
inline

Definition at line 117 of file TritonTool.cxx.

117 {
118 retryDelayMs *= (1 << attempt);
119 if (retryDelayMs > 0) {
120 std::this_thread::sleep_for(std::chrono::milliseconds(retryDelayMs));
121 }
122 }

Member Data Documentation

◆ ATLAS_THREAD_SAFE

std::atomic_flag m_initialized AthMessaging::ATLAS_THREAD_SAFE = ATOMIC_FLAG_INIT
mutableprivateinherited

Messaging initialized (initMessaging).

Definition at line 141 of file AthMessaging.h.

◆ m_imsg

std::atomic<IMessageSvc*> AthMessaging::m_imsg { nullptr }
mutableprivateinherited

MessageSvc pointer.

Definition at line 135 of file AthMessaging.h.

135{ nullptr };

◆ m_lvl

std::atomic<MSG::Level> AthMessaging::m_lvl { MSG::NIL }
mutableprivateinherited

Current logging level.

Definition at line 138 of file AthMessaging.h.

138{ MSG::NIL };

◆ m_msg_tls

boost::thread_specific_ptr<MsgStream> AthMessaging::m_msg_tls
mutableprivateinherited

MsgStream instance (a std::cout like with print-out levels).

Definition at line 132 of file AthMessaging.h.

◆ m_nm

std::string AthMessaging::m_nm
privateinherited

Message source name.

Definition at line 129 of file AthMessaging.h.

◆ m_options

std::unique_ptr<tc::InferOptions> AthInfer::TritonTool::Impl::m_options

Definition at line 235 of file TritonTool.cxx.

◆ m_parentAsyncAlg

const AthAsynchronousAlgorithm* AthInfer::TritonTool::Impl::m_parentAsyncAlg = nullptr

Definition at line 234 of file TritonTool.cxx.


The documentation for this struct was generated from the following file: