MCPcopy Create free account
hub / github.com/Oneflow-Inc/oneflow / CleanUselessMemBlockAndCheckValid

Method CleanUselessMemBlockAndCheckValid

oneflow/core/job/plan_util.cpp:475–563  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

473}
474
475void PlanUtil::CleanUselessMemBlockAndCheckValid(Plan* plan) {
476 HashMap<int64_t, ChunkProto> chunk_id2chunk;
477 HashMap<int64_t, MemBlockProto> mem_block_id2mem_block;
478 for (const auto& chunk : plan->block_chunk_list().chunk()) {
479 CHECK(chunk_id2chunk.emplace(chunk.chunk_id(), chunk).second);
480 }
481 for (const auto& mem_block : plan->block_chunk_list().mem_block()) {
482 CHECK(mem_block_id2mem_block.emplace(mem_block.mem_block_id(), mem_block).second);
483 }
484 plan->mutable_block_chunk_list()->clear_mem_block();
485
486 HashMap<int64_t, HashSet<int64_t>> chunk_id2job_ids;
487 HashMap<int64_t, HashSet<int64_t>> mem_block_id2job_ids;
488 for (const auto& pair : chunk_id2chunk) {
489 for (int64_t job_id : pair.second.job_id()) {
490 CHECK(chunk_id2job_ids[pair.first].insert(job_id).second);
491 }
492 }
493 for (const auto& pair : mem_block_id2mem_block) {
494 for (int64_t job_id : pair.second.job_id()) {
495 CHECK(mem_block_id2job_ids[pair.first].insert(job_id).second);
496 }
497 }
498
499 HashSet<int64_t> valid_mem_block_ids;
500 for (const TaskProto& task : plan->task()) {
501 for (const auto& pair : task.produced_regst_desc()) {
502 const RegstDescProto& regst = pair.second;
503 RtRegstDesc rt_regst(regst);
504 int64_t regst_size = rt_regst.TotalMainByteSize4AllRegst();
505 CHECK(mem_block_id2mem_block.find(regst.mem_block_id()) != mem_block_id2mem_block.end());
506 const MemBlockProto& mem_block = mem_block_id2mem_block.at(regst.mem_block_id());
507 CHECK_GE(mem_block.mem_size(), regst.mem_block_offset() + regst_size);
508 CHECK_EQ(task.machine_id(), mem_block.machine_id());
509 CHECK_EQ(mem_block.enable_reuse_mem(), regst.enable_reuse_mem());
510 CHECK(mem_block.mem_case() == regst.mem_case());
511 const auto& job_ids = mem_block_id2job_ids[regst.mem_block_id()];
512 CHECK(job_ids.find(task.job_id()) != job_ids.end());
513 valid_mem_block_ids.insert(regst.mem_block_id());
514
515 // separated_header
516 int64_t separated_header_mem_size = rt_regst.TotalSeparatedHeaderByteSize4AllRegst();
517 if (separated_header_mem_size > 0) {
518 int64_t header_block_id = regst.separated_header_mem_block_id();
519 CHECK_NE(header_block_id, -1);
520 CHECK(mem_block_id2mem_block.find(header_block_id) != mem_block_id2mem_block.end());
521 const MemBlockProto& header_mem_block = mem_block_id2mem_block.at(header_block_id);
522 CHECK_EQ(header_mem_block.mem_size(), separated_header_mem_size);
523 CHECK_EQ(task.machine_id(), header_mem_block.machine_id());
524 CHECK(header_mem_block.mem_case() == memory::GetPinnedHostMemoryCase(regst.mem_case()));
525 CHECK(header_mem_block.enable_reuse_mem() == false);
526 const auto& header_block_job_ids = mem_block_id2job_ids[header_block_id];
527 CHECK(header_block_job_ids.find(task.job_id()) != header_block_job_ids.end());
528 valid_mem_block_ids.insert(regst.separated_header_mem_block_id());
529 }
530 }
531 }
532

Callers

nothing calls this directly

Calls 14

GetPinnedHostMemoryCaseFunction · 0.85
mem_block_idMethod · 0.80
insertMethod · 0.80
findMethod · 0.80
mem_sizeMethod · 0.80
mem_block_offsetMethod · 0.80
job_idMethod · 0.45
endMethod · 0.45
atMethod · 0.45
machine_idMethod · 0.45

Tested by

no test coverage detected