| 277 | } |
| 278 | |
| 279 | void StepStatsCollector::BuildCostModel( |
| 280 | CostModelManager* cost_model_manager, |
| 281 | const std::unordered_map<string, const Graph*>& device_map) { |
| 282 | mutex_lock lock(mu_); |
| 283 | |
| 284 | if (!finalized_) { |
| 285 | FinalizeInternal(); |
| 286 | } |
| 287 | // Hardware stats for gpu are available under a fake device named |
| 288 | // "gpu:<id>/stream::all. |
| 289 | // Use them instead of regular stats whenever they're available to extract |
| 290 | // the execution stats of a particular node since they're more accurate. |
| 291 | // However hardware traces don't record memory usage, so we still have to |
| 292 | // rely on regular traces to track memory usage. |
| 293 | struct DeviceStats { |
| 294 | const DeviceStepStats* regular_stats; |
| 295 | const DeviceStepStats* hardware_stats; |
| 296 | }; |
| 297 | |
| 298 | std::unordered_map<StringPiece, DeviceStats, StringPieceHasher> |
| 299 | per_device_stats; |
| 300 | std::unordered_map<int, const DeviceStepStats*> gpu_hardware_stats; |
| 301 | |
| 302 | for (int i = 0; i < step_stats_->dev_stats_size(); ++i) { |
| 303 | const DeviceStepStats& device_stats = step_stats_->dev_stats(i); |
| 304 | const string& device_name = device_stats.device(); |
| 305 | const int gpu_id = ExtractGpuWithStreamAll(device_name); |
| 306 | if (gpu_id >= 0) { |
| 307 | // These are gpu hardware stats |
| 308 | gpu_hardware_stats.emplace(gpu_id, &device_stats); |
| 309 | } else { |
| 310 | // These are regular stats. |
| 311 | per_device_stats.emplace(device_name, |
| 312 | DeviceStats{&device_stats, nullptr}); |
| 313 | } |
| 314 | } |
| 315 | |
| 316 | for (auto& itr : per_device_stats) { |
| 317 | const StringPiece device_name = itr.first; |
| 318 | const int gpu_id = ExtractGpuWithoutStream(string(device_name)); |
| 319 | if (gpu_id >= 0) { |
| 320 | // Reference the gpu hardware stats in addition to the regular stats |
| 321 | // for this gpu device if they're available. |
| 322 | if (gpu_hardware_stats.find(gpu_id) != gpu_hardware_stats.end()) { |
| 323 | itr.second.hardware_stats = gpu_hardware_stats.find(gpu_id)->second; |
| 324 | } |
| 325 | } |
| 326 | } |
| 327 | |
| 328 | for (auto itr : device_map) { |
| 329 | const StringPiece device = itr.first; |
| 330 | if (per_device_stats.find(device) == per_device_stats.end()) { |
| 331 | continue; |
| 332 | } |
| 333 | |
| 334 | const Graph* graph = itr.second; |
| 335 | CostModel* cm = cost_model_manager->FindOrCreateCostModel(graph); |
| 336 | cm->IncrementUpdateTimes(); |
no test coverage detected