| 105 | c_state_.resize(c_size, 0.0f); |
| 106 | } |
| 107 | std::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 | } |