-
Notifications
You must be signed in to change notification settings - Fork 84
Conversation
Description
|
| Relevant files | |
|---|---|
| Enhancement |
id_model.cpp
Implement almost exact split mapping functionality csrc/id_model/id_model.cpp equivalence logic loop_domain_scheduler.cpp
Add ValGraph parameter support to scheduler csrc/scheduler/tools/loop_domain_scheduler.cpp graphs id_model.h
Add function declaration for split mapping csrc/id_model/id_model.h
loop_domain_scheduler.h
Update scheduler interface for graph parameter csrc/scheduler/tools/loop_domain_scheduler.h parameter |
| Tests |
test_id_model.cpp
Add comprehensive tests for almost exact split functionalitytests/cpp/test_id_model.cpp AlmostExactSplitGraph5 |
PR Reviewer Guide
Here are some key observations to aid the review process:
Debug Output
The new mapAlmostExactSplits function contains multiple std::cerr debugging outputs (lines 1474, 1475, 1518, 1519, 1532, 1533, 1539, 1557-1559) that should be removed or made conditional for production code.
ValGraph mapAlmostExactSplits(const ValGraph& graph) { auto new_graph = graph; // vg: I0 auto get_l1r2_splits = [&new_graph]( const ValGroup& vg) -> std::vector<std::pair<ExprGroup, ExprGroup>> { std::vector<std::pair<ExprGroup, ExprGroup>> l1_r2_splits; if (!new_graph.hasUses(vg)) { return {}; } for (const ExprGroup& use_of_vg : new_graph.getUses(vg)) { auto split_of_vg = dynamic_cast<Split*>(use_of_vg->front()); if (split_of_vg == nullptr) { continue; } // mn const ValGroup& inner_group = new_graph.toGroup(split_of_vg->inner()); if (!new_graph.hasUses(inner_group)) { return {}; } for (const ExprGroup& use_of_inner_group : new_graph.getUses(inner_group)) { auto split_of_inner_group = dynamic_cast<Split*>(use_of_inner_group->front()); if (split_of_inner_group == nullptr) { continue; } // This split needs to be divisible auto extent = split_of_inner_group->in()->extent(); auto factor = split_of_inner_group->factor(); if (extent->isConstScalar() && factor->isConstScalar() && (extent->evaluate().as<int64_t>() % factor->evaluate().as<int64_t>() != 0)) { continue; } l1_r2_splits.emplace_back(use_of_vg, use_of_inner_group); std::cerr << "L1R2 found: " << split_of_vg->toString() << split_of_inner_group->toString(); } } return l1_r2_splits; }; auto get_matching_l2r1_splits = [&new_graph]( const ValGroup& vg, const std::pair<ExprGroup, ExprGroup>& l1_r2) -> std::optional<std::pair<ExprGroup, ExprGroup>> { auto m = l1_r2.second->front()->as<Split>()->outer()->extent(); auto n = l1_r2.second->front()->as<Split>()->inner()->extent(); for (const ExprGroup& use_of_vg : new_graph.getUses(vg)) { auto split_of_vg = dynamic_cast<Split*>(use_of_vg->front()); if (split_of_vg == nullptr) { continue; } if (!split_of_vg->inner()->extent()->sameAs(n)) { continue; } // I0/n const ValGroup& outer_group = new_graph.toGroup(split_of_vg->outer()); if (!new_graph.hasUses(outer_group)) { return {}; } for (const ExprGroup& use_of_outer_group : new_graph.getUses(outer_group)) { auto split_of_outer_group = dynamic_cast<Split*>(use_of_outer_group->front()); if (split_of_outer_group == nullptr) { continue; } if (!split_of_outer_group->inner()->extent()->sameAs(m)) { continue; } std::cerr << "Matching L2R1 found: " << split_of_vg->toString() << split_of_outer_group->toString(); return std::make_pair(use_of_vg, use_of_outer_group); } } return std::nullopt; }; std::vector<std::pair<ValGroup, ValGroup>> groups_to_map; for (const ValGroup& vg : new_graph.disjointValSets().disjointSets()) { const auto all_l1r2_splits = get_l1r2_splits(vg); for (const auto& l1r2 : all_l1r2_splits) { std::cerr << "L1R2: " << l1r2.first->front()->toString() << l1r2.second->front()->toString(); auto l2r1 = get_matching_l2r1_splits(vg, l1r2); if (!l2r1.has_value()) { continue; } std::cerr << "Found\n"; auto l1r2_first_outputs = new_graph.outputGroups(l1r2.first); auto l1r2_second_outputs = new_graph.outputGroups(l1r2.second); auto l2r1_first_outputs = new_graph.outputGroups(l2r1->first); auto l2r1_second_outputs = new_graph.outputGroups(l2r1->second); groups_to_map.emplace_back( l1r2_first_outputs.at(0), l2r1_second_outputs.at(0)); groups_to_map.emplace_back( l1r2_second_outputs.at(0), l2r1_second_outputs.at(1)); groups_to_map.emplace_back( l1r2_second_outputs.at(1), l2r1_first_outputs.at(1)); } } for (const auto& [vg1, vg2] : groups_to_map) { std::cerr << "Mapping " << nvfuser::toString(vg1) << ", " << vg1->front()->toString() << " and " << nvfuser::toString(vg2) << ", " << vg2->front()->toString() << "\n"; new_graph.mapVals(vg1->front(), vg2->front()); } return new_graph; }
Error Handling
The function has several early returns with empty vectors without proper error handling or logging when expected split patterns are not found, which could make debugging difficult.
if (!new_graph.hasUses(vg)) { return {}; } for (const ExprGroup& use_of_vg : new_graph.getUses(vg)) { auto split_of_vg = dynamic_cast<Split*>(use_of_vg->front()); if (split_of_vg == nullptr) { continue; } // mn const ValGroup& inner_group = new_graph.toGroup(split_of_vg->inner()); if (!new_graph.hasUses(inner_group)) { return {}; } for (const ExprGroup& use_of_inner_group : new_graph.getUses(inner_group)) { auto split_of_inner_group = dynamic_cast<Split*>(use_of_inner_group->front()); if (split_of_inner_group == nullptr) { continue; } // This split needs to be divisible auto extent = split_of_inner_group->in()->extent(); auto factor = split_of_inner_group->factor(); if (extent->isConstScalar() && factor->isConstScalar() && (extent->evaluate().as<int64_t>() % factor->evaluate().as<int64_t>() != 0)) { continue; } l1_r2_splits.emplace_back(use_of_vg, use_of_inner_group); std::cerr << "L1R2 found: " << split_of_vg->toString() << split_of_inner_group->toString(); } } return l1_r2_splits; }; auto get_matching_l2r1_splits = [&new_graph]( const ValGroup& vg, const std::pair<ExprGroup, ExprGroup>& l1_r2) -> std::optional<std::pair<ExprGroup, ExprGroup>> { auto m = l1_r2.second->front()->as<Split>()->outer()->extent(); auto n = l1_r2.second->front()->as<Split>()->inner()->extent(); for (const ExprGroup& use_of_vg : new_graph.getUses(vg)) { auto split_of_vg = dynamic_cast<Split*>(use_of_vg->front()); if (split_of_vg == nullptr) { continue; } if (!split_of_vg->inner()->extent()->sameAs(n)) { continue; } // I0/n const ValGroup& outer_group = new_graph.toGroup(split_of_vg->outer()); if (!new_graph.hasUses(outer_group)) { return {}; } for (const ExprGroup& use_of_outer_group : new_graph.getUses(outer_group)) { auto split_of_outer_group = dynamic_cast<Split*>(use_of_outer_group->front()); if (split_of_outer_group == nullptr) { continue; } if (!split_of_outer_group->inner()->extent()->sameAs(m)) { continue; } std::cerr << "Matching L2R1 found: " << split_of_vg->toString() << split_of_outer_group->toString(); return std::make_pair(use_of_vg, use_of_outer_group); } } return std::nullopt;
Memory Management
The constructor now accepts an optional graph pointer but doesn't document ownership semantics or lifetime requirements, which could lead to dangling pointer issues.
LoopDomainScheduler( std::vector<IterDomain*> ref_loop_dom, bool update_loop_domain_only = false, const ValGraph* scheduling_graph = nullptr) : ref_loop_dom_(std::move(ref_loop_dom)), update_loop_domain_only_(update_loop_domain_only), graph_(scheduling_graph) { NVF_ERROR(!ref_loop_dom_.empty()); if (graph_ == nullptr) { Fusion* fusion = ref_loop_dom_.front()->fusion(); id_model_ = std::make_unique<IdModel>(fusion, /*build_graphs=*/false); id_model_->buildExactGraph(); graph_ = &(id_model_->idGraph(IdMappingMode::EXACT)); }
No description provided.