Skip to content

Navigation Menu

Sign in
Sign up

[IdModel] Equivalence of Split-Split - #5984

Draft
wujingyue wants to merge 4 commits into
main from
almost_exact_split
Draft

[IdModel] Equivalence of Split-Split #5984
wujingyue wants to merge 4 commits into
main from
almost_exact_split

Conversation

@wujingyue

@wujingyue wujingyue commented Feb 19, 2026

Copy link
Copy Markdown
Collaborator

No description provided.

Naoya Maruyama added 4 commits February 19, 2026 10:08

Copy link
Copy Markdown

Description

  • Implement mapAlmostExactSplits function to handle split transformation equivalence

  • Add support for custom ValGraph in loop domain scheduling

  • Create 5 test cases validating almost exact split mapping patterns

  • Enable flexible graph-based scheduling with optional graph parameter

Changes walkthrough

Relevant files
Enhancement
id_model.cpp
Implement almost exact split mapping functionality

csrc/id_model/id_model.cpp

  • Remove unused includes for trivial_broadcast.h and val_graph_visitor.h
  • Add mapAlmostExactSplits function implementing split pattern
    equivalence logic
  • Implement L1R2 and L2R1 split detection and mapping algorithm
  • +139/-2
    loop_domain_scheduler.cpp
    Add ValGraph parameter support to scheduler

    csrc/scheduler/tools/loop_domain_scheduler.cpp

  • Add optional ValGraph parameter to LoopDomainScheduler constructor
  • Modify scheduleLoopDomainsLike to accept custom ValGraph
  • Enable using pre-built graphs instead of always creating new exact
    graphs
  • +17/-9
    id_model.h
    Add function declaration for split mapping

    csrc/id_model/id_model.h

    • Add declaration for mapAlmostExactSplits function
    +2/-0
    loop_domain_scheduler.h
    Update scheduler interface for graph parameter

    csrc/scheduler/tools/loop_domain_scheduler.h

  • Update scheduleLoopDomainsLike signature with optional ValGraph
    parameter
  • +2/-1
    Tests
    test_id_model.cpp
    Add comprehensive tests for almost exact split functionality

    tests/cpp/test_id_model.cpp

  • Add 5 new test cases: AlmostExactSplitGraph1 through
    AlmostExactSplitGraph5
  • Remove unused includes for fstream and ir/graphviz.h
  • Add scheduler_tools include for loop domain scheduling
  • Test various split/reshape/merge patterns with almost exact mapping
  • +235/-10

    PR Reviewer Guide

    Here are some key observations to aid the review process:

    🧪 PR contains tests
    Recommended focus areas for review
    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));
     }
    

    Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

    Reviewers

    No reviews

    Assignees

    No one assigned

    Labels

    None yet

    Projects

    None yet

    Milestone

    No milestone

    Development

    Successfully merging this pull request may close these issues.

    1 participant

    AltStyle によって変換されたページ (->オリジナル) /