MCPcopy Create free account
hub / github.com/TeleHuman/TextOp / predict

Method predict

TextOpDeploy/src/textop_ctrl/src/onnx_policy.cpp:107–185  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

105 c_state_.resize(c_size, 0.0f);
106}
107std::vector<float> ONNXPolicy::predict(const std::vector<float>& observation)
108{
109 try
110 {
111 size_t num_inputs = session_->GetInputCount();
112
113 // Prepare input tensors
114 std::vector<Ort::Value> input_tensors;
115
116 // Always add observation as first input
117 std::vector<int64_t> obs_shape = {1, static_cast<int64_t>(observation.size())};
118 input_tensors.emplace_back(Ort::Value::CreateTensor<float>(
119 Ort::MemoryInfo::CreateCpu(OrtArenaAllocator, OrtMemTypeDefault),
120 const_cast<float*>(observation.data()), observation.size(), obs_shape.data(),
121 obs_shape.size()));
122
123 // For multi-input models (LSTM), add hidden and cell states
124 if (num_inputs >= 3)
125 {
126 // Hidden state
127 input_tensors.emplace_back(Ort::Value::CreateTensor<float>(
128 Ort::MemoryInfo::CreateCpu(OrtArenaAllocator, OrtMemTypeDefault), h_state_.data(),
129 h_state_.size(), h_shape_.data(), h_shape_.size()));
130
131 // Cell state
132 input_tensors.emplace_back(Ort::Value::CreateTensor<float>(
133 Ort::MemoryInfo::CreateCpu(OrtArenaAllocator, OrtMemTypeDefault), c_state_.data(),
134 c_state_.size(), c_shape_.data(), c_shape_.size()));
135 }
136
137 // Run inference
138 auto output_tensors =
139 session_->Run(Ort::RunOptions{nullptr}, input_names_.data(), input_tensors.data(),
140 input_tensors.size(), output_names_.data(), output_names_.size());
141
142 // Extract output
143 float* output_data = output_tensors[0].GetTensorMutableData<float>();
144 auto output_shape = output_tensors[0].GetTensorTypeAndShapeInfo().GetShape();
145 size_t output_size = 1;
146 for (auto dim : output_shape)
147 {
148 output_size *= static_cast<size_t>(dim);
149 }
150
151 std::vector<float> result(output_data, output_data + output_size);
152
153 // Update memory states for LSTM models
154 if (num_inputs >= 3 && output_tensors.size() >= 3)
155 {
156 float* new_h_data = output_tensors[1].GetTensorMutableData<float>();
157 float* new_c_data = output_tensors[2].GetTensorMutableData<float>();
158
159 std::copy(new_h_data, new_h_data + h_state_.size(), h_state_.begin());
160 std::copy(new_c_data, new_c_data + c_state_.size(), c_state_.begin());
161 }
162
163 return result;
164 }

Callers 1

ControlMethod · 0.80

Calls

no outgoing calls

Tested by

no test coverage detected