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 (const int retryDelayMs) 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 216 of file TritonTool.cxx.

218 {
219
220 const uint8_t* rawData = nullptr;
221 size_t size = 0;
222
223 // Get access to the buffer holding raw results of specified output returned
224 // by the server. Note: the buffer is owned by InferResult instance. Users
225 // can copy out the data if required to extend the lifetime.
226 TRITON_CHECK(result.RawData(name, &rawData, &size));
227
228 outputVec.resize(size / sizeof(T));
229 std::memcpy(outputVec.data(), rawData, size);
230 return StatusCode::SUCCESS;
231 }
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)

◆ 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 188 of file TritonTool.cxx.

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

◆ 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 123 of file TritonTool.cxx.

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

Definition at line 117 of file TritonTool.cxx.

117 {
118 if (retryDelayMs > 0) {
119 std::this_thread::sleep_for(std::chrono::milliseconds(retryDelayMs));
120 }
121 }

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 234 of file TritonTool.cxx.

◆ m_parentAsyncAlg

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

Definition at line 233 of file TritonTool.cxx.


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