Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 42 additions & 0 deletions include/tvm/node/structural_equal.h
Original file line numberDiff line numberDiff line change
Expand Up@@ -324,5 +324,47 @@ class SEqualReducer {
bool map_free_vars_ = false;
};

/*! \brief The default handler for equality testing.
*
* Users can derive from this class and override the DispatchSEqualReduce method,
* to customize equality testing.
*/
class SEqualHandlerDefault : public SEqualReducer::Handler {
public:
SEqualHandlerDefault(bool assert_mode, Optional<ObjectPathPair>* first_mismatch);
virtual ~SEqualHandlerDefault();

bool SEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) override;
void DeferFail(const ObjectPathPair& mismatch_paths) override;
ObjectRef MapLhsToRhs(const ObjectRef& lhs) override;
void MarkGraphNode() override;

/*!
* \brief The entry point for equality testing
* \param lhs The left operand.
* \param rhs The right operand.
* \param map_free_vars Whether or not to remap variables if possible.
* \return The equality result.
*/
virtual bool Equal(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars);

protected:
/*!
* \brief The dispatcher for equality testing of intermediate objects
* \param lhs The left operand.
* \param rhs The right operand.
* \param map_free_vars Whether or not to remap variables if possible.
* \param current_paths Optional paths to `lhs` and `rhs` objects, for error traceability.
* \return The equality result.
*/
virtual bool DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths);

private:
class Impl;
Impl* impl;
};

} // namespace tvm
#endif // TVM_NODE_STRUCTURAL_EQUAL_H_
38 changes: 38 additions & 0 deletions include/tvm/node/structural_hash.h
Original file line numberDiff line numberDiff line change
Expand Up@@ -200,6 +200,44 @@ class SHashReducer {
bool map_free_vars_;
};

/*! \brief The default handler for hash key computation
*
* Users can derive from this class and override the DispatchSHash method,
* to customize hashing.
*/
class SHashHandlerDefault : public SHashReducer::Handler {
public:
SHashHandlerDefault();
virtual ~SHashHandlerDefault();

void SHashReduceHashedValue(size_t hashed_value) override;
void SHashReduce(const ObjectRef& key, bool map_free_vars) override;
void SHashReduceFreeVar(const runtime::Object* var, bool map_free_vars) override;
bool LookupHashedValue(const ObjectRef& key, size_t* hashed_value) override;
void MarkGraphNode() override;

/*!
* \brief The entry point for hashing
* \param object The object to be hashed.
* \param map_free_vars Whether or not to remap variables if possible.
* \return The hash result.
*/
virtual size_t Hash(const ObjectRef& object, bool map_free_vars);

protected:
/*!
* \brief The dispatcher for hashing of intermediate objects
* \param object An intermediate object to be hashed.
* \param map_free_vars Whether or not to remap variables if possible.
* \return The hash result.
*/
virtual void DispatchSHash(const ObjectRef& object, bool map_free_vars);

private:
class Impl;
Impl* impl;
};

class SEqualReducer;
struct NDArrayContainerTrait {
static constexpr const std::nullptr_t VisitAttrs = nullptr;
Expand Down
109 changes: 73 additions & 36 deletions src/node/structural_equal.cc
Original file line numberDiff line numberDiff line change
Expand Up@@ -198,13 +198,13 @@ bool SEqualReducer::ObjectAttrsEqual(const ObjectRef& lhs, const ObjectRef& rhs,
* The order of SEqual being called is the same as the order as if we
* eagerly do recursive calls in SEqualReduce.
*/
class RemapVarSEqualHandler : public SEqualReducer::Handler {
class SEqualHandlerDefault::Impl {
public:
explicit RemapVarSEqualHandler(bool assert_mode, Optional<ObjectPathPair>* first_mismatch)
: assert_mode_(assert_mode), first_mismatch_(first_mismatch) {}
Impl(SEqualHandlerDefault* parent, bool assert_mode, Optional<ObjectPathPair>* first_mismatch)
: parent_(parent), assert_mode_(assert_mode), first_mismatch_(first_mismatch) {}

bool SEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) final {
const Optional<ObjectPathPair>& current_paths) {
// We cannot use check lhs.same_as(rhs) to check equality.
// if we choose to enable var remapping.
//
Expand DownExpand Up@@ -239,17 +239,17 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
return CheckResult(run(), lhs, rhs, current_paths);
}

void DeferFail(const ObjectPathPair& mismatch_paths) final {
void DeferFail(const ObjectPathPair& mismatch_paths) {
pending_tasks_.emplace_back(Task::ForceFailTag{}, mismatch_paths);
}

void MarkGraphNode() final {
void MarkGraphNode() {
// need to push to pending tasks in this case
ICHECK(!allow_push_to_stack_ && !task_stack_.empty());
task_stack_.back().graph_equal = true;
}

ObjectRef MapLhsToRhs(const ObjectRef& lhs) final {
ObjectRef MapLhsToRhs(const ObjectRef& lhs) {
auto it = equal_map_lhs_.find(lhs);
if (it != equal_map_lhs_.end()) return it->second;
return ObjectRef(nullptr);
Expand DownExpand Up@@ -279,7 +279,35 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
return RunTasks();
}

// The default equal as registered in the structural equal vtable.
bool DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
auto compute = [=]() {
ICHECK(lhs.defined() && rhs.defined() && lhs->type_index() == rhs->type_index());
// skip entries that already have equality maps.
auto it = equal_map_lhs_.find(lhs);
if (it != equal_map_lhs_.end()) {
return it->second.same_as(rhs);
}
if (equal_map_rhs_.count(rhs)) return false;

SEqualReducer reducer = GetReducer(lhs, rhs, map_free_vars, current_paths);
return vtable_->SEqualReduce(lhs.get(), rhs.get(), reducer);
};
return CheckResult(compute(), lhs, rhs, current_paths);
}

protected:
SEqualReducer GetReducer(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
if (!IsPathTracingEnabled()) {
return SEqualReducer(parent_, nullptr, map_free_vars);
} else {
PathTracingData tracing_data = {current_paths.value(), lhs, rhs, first_mismatch_};
return SEqualReducer(parent_, &tracing_data, map_free_vars);
}
}

// Check the result.
bool CheckResult(bool result, const ObjectRef& lhs, const ObjectRef& rhs,
const Optional<ObjectPathPair>& current_paths) {
Expand DownExpand Up@@ -335,7 +363,8 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
// which populates the pending tasks.
ICHECK_EQ(pending_tasks_.size(), 0U);
allow_push_to_stack_ = false;
if (!DispatchSEqualReduce(entry.lhs, entry.rhs, entry.map_free_vars, entry.current_paths))
if (!parent_->DispatchSEqualReduce(entry.lhs, entry.rhs, entry.map_free_vars,
entry.current_paths))
return false;
allow_push_to_stack_ = true;
// Push pending tasks in reverse order, so earlier tasks get to
Expand All@@ -349,31 +378,6 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
return true;
}

// The default equal as registered in the structural equal vtable.
bool DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
auto compute = [=]() {
ICHECK(lhs.defined() && rhs.defined() && lhs->type_index() == rhs->type_index());
// skip entries that already have equality maps.
auto it = equal_map_lhs_.find(lhs);
if (it != equal_map_lhs_.end()) {
return it->second.same_as(rhs);
}
if (equal_map_rhs_.count(rhs)) return false;

// Run reduce check for free nodes.
if (!IsPathTracingEnabled()) {
return vtable_->SEqualReduce(lhs.get(), rhs.get(),
SEqualReducer(this, nullptr, map_free_vars));
} else {
PathTracingData tracing_data = {current_paths.value(), lhs, rhs, first_mismatch_};
return vtable_->SEqualReduce(lhs.get(), rhs.get(),
SEqualReducer(this, &tracing_data, map_free_vars));
}
};
return CheckResult(compute(), lhs, rhs, current_paths);
}

private:
/*! \brief Pending reduce tasks. */
struct Task {
Expand DownExpand Up@@ -407,6 +411,8 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {

bool IsPathTracingEnabled() const { return first_mismatch_ != nullptr; }

// The owner of this impl
SEqualHandlerDefault* parent_;
// list of pending tasks to be pushed to the stack.
std::vector<Task> pending_tasks_;
// Internal task stack to executed the task.
Expand All@@ -425,22 +431,53 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
std::unordered_map<ObjectRef, ObjectRef, ObjectPtrHash, ObjectPtrEqual> equal_map_rhs_;
};

SEqualHandlerDefault::SEqualHandlerDefault(bool assert_mode,
Optional<ObjectPathPair>* first_mismatch) {
impl = new Impl(this, assert_mode, first_mismatch);
}

SEqualHandlerDefault::~SEqualHandlerDefault() { delete impl; }

bool SEqualHandlerDefault::SEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs,
bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
return impl->SEqualReduce(lhs, rhs, map_free_vars, current_paths);
}

void SEqualHandlerDefault::DeferFail(const ObjectPathPair& mismatch_paths) {
impl->DeferFail(mismatch_paths);
}

ObjectRef SEqualHandlerDefault::MapLhsToRhs(const ObjectRef& lhs) { return impl->MapLhsToRhs(lhs); }

void SEqualHandlerDefault::MarkGraphNode() { impl->MarkGraphNode(); }

bool SEqualHandlerDefault::Equal(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars) {
return impl->Equal(lhs, rhs, map_free_vars);
}

bool SEqualHandlerDefault::DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs,
bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
return impl->DispatchSEqualReduce(lhs, rhs, map_free_vars, current_paths);
}

TVM_REGISTER_GLOBAL("node.StructuralEqual")
.set_body_typed([](const ObjectRef& lhs, const ObjectRef& rhs, bool assert_mode,
bool map_free_vars) {
return RemapVarSEqualHandler(assert_mode, nullptr).Equal(lhs, rhs, map_free_vars);
return SEqualHandlerDefault(assert_mode, nullptr).Equal(lhs, rhs, map_free_vars);
});

TVM_REGISTER_GLOBAL("node.GetFirstStructuralMismatch")
.set_body_typed([](const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars) {
Optional<ObjectPathPair> first_mismatch;
bool equal = RemapVarSEqualHandler(false, &first_mismatch).Equal(lhs, rhs, map_free_vars);
bool equal = SEqualHandlerDefault(false, &first_mismatch).Equal(lhs, rhs, map_free_vars);
ICHECK(equal == !first_mismatch.defined());
return first_mismatch;
});

bool StructuralEqual::operator()(const ObjectRef& lhs, const ObjectRef& rhs) const {
return RemapVarSEqualHandler(false, nullptr).Equal(lhs, rhs, false);
return SEqualHandlerDefault(false, nullptr).Equal(lhs, rhs, false);
}

} // namespace tvm
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Add copy buttons to all
 blocks
(function() {
function addCopyButtons() {
document.querySelectorAll('pre code').forEach(function(codeBlock) {
if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;
codeBlock.parentElement.setAttribute('data-copy-added', 'true');
var btn = document.createElement('button');
btn.textContent = 'Copy';
btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';
btn.onmouseover = function() { this.style.opacity = '1'; };
btn.onmouseout = function() { this.style.opacity = '0.7'; };
btn.onclick = function() {
navigator.clipboard.writeText(codeBlock.textContent).then(function() {
btn.textContent = 'Copied!';
setTimeout(function() { btn.textContent = 'Copy'; }, 1500);
});
};
codeBlock.parentElement.style.position = 'relative';
codeBlock.parentElement.appendChild(btn);
});
}
addCopyButtons();
// Re-run on dynamic content
var observer = new MutationObserver(addCopyButtons);
observer.observe(document.body, { childList: true, subtree: true });
})();
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 42 additions & 0 deletions include/tvm/node/structural_equal.h
Original file line numberDiff line numberDiff line change
Expand Up@@ -324,5 +324,47 @@ class SEqualReducer {
bool map_free_vars_ = false;
};

/*! \brief The default handler for equality testing.
*
* Users can derive from this class and override the DispatchSEqualReduce method,
* to customize equality testing.
*/
class SEqualHandlerDefault : public SEqualReducer::Handler {
public:
SEqualHandlerDefault(bool assert_mode, Optional<ObjectPathPair>* first_mismatch);
virtual ~SEqualHandlerDefault();

bool SEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) override;
void DeferFail(const ObjectPathPair& mismatch_paths) override;
ObjectRef MapLhsToRhs(const ObjectRef& lhs) override;
void MarkGraphNode() override;

/*!
* \brief The entry point for equality testing
* \param lhs The left operand.
* \param rhs The right operand.
* \param map_free_vars Whether or not to remap variables if possible.
* \return The equality result.
*/
virtual bool Equal(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars);

protected:
/*!
* \brief The dispatcher for equality testing of intermediate objects
* \param lhs The left operand.
* \param rhs The right operand.
* \param map_free_vars Whether or not to remap variables if possible.
* \param current_paths Optional paths to `lhs` and `rhs` objects, for error traceability.
* \return The equality result.
*/
virtual bool DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths);

private:
class Impl;
Impl* impl;
};

} // namespace tvm
#endif // TVM_NODE_STRUCTURAL_EQUAL_H_
38 changes: 38 additions & 0 deletions include/tvm/node/structural_hash.h
Original file line numberDiff line numberDiff line change
Expand Up@@ -200,6 +200,44 @@ class SHashReducer {
bool map_free_vars_;
};

/*! \brief The default handler for hash key computation
*
* Users can derive from this class and override the DispatchSHash method,
* to customize hashing.
*/
class SHashHandlerDefault : public SHashReducer::Handler {
public:
SHashHandlerDefault();
virtual ~SHashHandlerDefault();

void SHashReduceHashedValue(size_t hashed_value) override;
void SHashReduce(const ObjectRef& key, bool map_free_vars) override;
void SHashReduceFreeVar(const runtime::Object* var, bool map_free_vars) override;
bool LookupHashedValue(const ObjectRef& key, size_t* hashed_value) override;
void MarkGraphNode() override;

/*!
* \brief The entry point for hashing
* \param object The object to be hashed.
* \param map_free_vars Whether or not to remap variables if possible.
* \return The hash result.
*/
virtual size_t Hash(const ObjectRef& object, bool map_free_vars);

protected:
/*!
* \brief The dispatcher for hashing of intermediate objects
* \param object An intermediate object to be hashed.
* \param map_free_vars Whether or not to remap variables if possible.
* \return The hash result.
*/
virtual void DispatchSHash(const ObjectRef& object, bool map_free_vars);

private:
class Impl;
Impl* impl;
};

class SEqualReducer;
struct NDArrayContainerTrait {
static constexpr const std::nullptr_t VisitAttrs = nullptr;
Expand Down
109 changes: 73 additions & 36 deletions src/node/structural_equal.cc
Original file line numberDiff line numberDiff line change
Expand Up@@ -198,13 +198,13 @@ bool SEqualReducer::ObjectAttrsEqual(const ObjectRef& lhs, const ObjectRef& rhs,
* The order of SEqual being called is the same as the order as if we
* eagerly do recursive calls in SEqualReduce.
*/
class RemapVarSEqualHandler : public SEqualReducer::Handler {
class SEqualHandlerDefault::Impl {
public:
explicit RemapVarSEqualHandler(bool assert_mode, Optional<ObjectPathPair>* first_mismatch)
: assert_mode_(assert_mode), first_mismatch_(first_mismatch) {}
Impl(SEqualHandlerDefault* parent, bool assert_mode, Optional<ObjectPathPair>* first_mismatch)
: parent_(parent), assert_mode_(assert_mode), first_mismatch_(first_mismatch) {}

bool SEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) final {
const Optional<ObjectPathPair>& current_paths) {
// We cannot use check lhs.same_as(rhs) to check equality.
// if we choose to enable var remapping.
//
Expand DownExpand Up@@ -239,17 +239,17 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
return CheckResult(run(), lhs, rhs, current_paths);
}

void DeferFail(const ObjectPathPair& mismatch_paths) final {
void DeferFail(const ObjectPathPair& mismatch_paths) {
pending_tasks_.emplace_back(Task::ForceFailTag{}, mismatch_paths);
}

void MarkGraphNode() final {
void MarkGraphNode() {
// need to push to pending tasks in this case
ICHECK(!allow_push_to_stack_ && !task_stack_.empty());
task_stack_.back().graph_equal = true;
}

ObjectRef MapLhsToRhs(const ObjectRef& lhs) final {
ObjectRef MapLhsToRhs(const ObjectRef& lhs) {
auto it = equal_map_lhs_.find(lhs);
if (it != equal_map_lhs_.end()) return it->second;
return ObjectRef(nullptr);
Expand DownExpand Up@@ -279,7 +279,35 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
return RunTasks();
}

// The default equal as registered in the structural equal vtable.
bool DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
auto compute = [=]() {
ICHECK(lhs.defined() && rhs.defined() && lhs->type_index() == rhs->type_index());
// skip entries that already have equality maps.
auto it = equal_map_lhs_.find(lhs);
if (it != equal_map_lhs_.end()) {
return it->second.same_as(rhs);
}
if (equal_map_rhs_.count(rhs)) return false;

SEqualReducer reducer = GetReducer(lhs, rhs, map_free_vars, current_paths);
return vtable_->SEqualReduce(lhs.get(), rhs.get(), reducer);
};
return CheckResult(compute(), lhs, rhs, current_paths);
}

protected:
SEqualReducer GetReducer(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
if (!IsPathTracingEnabled()) {
return SEqualReducer(parent_, nullptr, map_free_vars);
} else {
PathTracingData tracing_data = {current_paths.value(), lhs, rhs, first_mismatch_};
return SEqualReducer(parent_, &tracing_data, map_free_vars);
}
}

// Check the result.
bool CheckResult(bool result, const ObjectRef& lhs, const ObjectRef& rhs,
const Optional<ObjectPathPair>& current_paths) {
Expand DownExpand Up@@ -335,7 +363,8 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
// which populates the pending tasks.
ICHECK_EQ(pending_tasks_.size(), 0U);
allow_push_to_stack_ = false;
if (!DispatchSEqualReduce(entry.lhs, entry.rhs, entry.map_free_vars, entry.current_paths))
if (!parent_->DispatchSEqualReduce(entry.lhs, entry.rhs, entry.map_free_vars,
entry.current_paths))
return false;
allow_push_to_stack_ = true;
// Push pending tasks in reverse order, so earlier tasks get to
Expand All@@ -349,31 +378,6 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
return true;
}

// The default equal as registered in the structural equal vtable.
bool DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
auto compute = [=]() {
ICHECK(lhs.defined() && rhs.defined() && lhs->type_index() == rhs->type_index());
// skip entries that already have equality maps.
auto it = equal_map_lhs_.find(lhs);
if (it != equal_map_lhs_.end()) {
return it->second.same_as(rhs);
}
if (equal_map_rhs_.count(rhs)) return false;

// Run reduce check for free nodes.
if (!IsPathTracingEnabled()) {
return vtable_->SEqualReduce(lhs.get(), rhs.get(),
SEqualReducer(this, nullptr, map_free_vars));
} else {
PathTracingData tracing_data = {current_paths.value(), lhs, rhs, first_mismatch_};
return vtable_->SEqualReduce(lhs.get(), rhs.get(),
SEqualReducer(this, &tracing_data, map_free_vars));
}
};
return CheckResult(compute(), lhs, rhs, current_paths);
}

private:
/*! \brief Pending reduce tasks. */
struct Task {
Expand DownExpand Up@@ -407,6 +411,8 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {

bool IsPathTracingEnabled() const { return first_mismatch_ != nullptr; }

// The owner of this impl
SEqualHandlerDefault* parent_;
// list of pending tasks to be pushed to the stack.
std::vector<Task> pending_tasks_;
// Internal task stack to executed the task.
Expand All@@ -425,22 +431,53 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
std::unordered_map<ObjectRef, ObjectRef, ObjectPtrHash, ObjectPtrEqual> equal_map_rhs_;
};

SEqualHandlerDefault::SEqualHandlerDefault(bool assert_mode,
Optional<ObjectPathPair>* first_mismatch) {
impl = new Impl(this, assert_mode, first_mismatch);
}

SEqualHandlerDefault::~SEqualHandlerDefault() { delete impl; }

bool SEqualHandlerDefault::SEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs,
bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
return impl->SEqualReduce(lhs, rhs, map_free_vars, current_paths);
}

void SEqualHandlerDefault::DeferFail(const ObjectPathPair& mismatch_paths) {
impl->DeferFail(mismatch_paths);
}

ObjectRef SEqualHandlerDefault::MapLhsToRhs(const ObjectRef& lhs) { return impl->MapLhsToRhs(lhs); }

void SEqualHandlerDefault::MarkGraphNode() { impl->MarkGraphNode(); }

bool SEqualHandlerDefault::Equal(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars) {
return impl->Equal(lhs, rhs, map_free_vars);
}

bool SEqualHandlerDefault::DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs,
bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
return impl->DispatchSEqualReduce(lhs, rhs, map_free_vars, current_paths);
}

TVM_REGISTER_GLOBAL("node.StructuralEqual")
.set_body_typed([](const ObjectRef& lhs, const ObjectRef& rhs, bool assert_mode,
bool map_free_vars) {
return RemapVarSEqualHandler(assert_mode, nullptr).Equal(lhs, rhs, map_free_vars);
return SEqualHandlerDefault(assert_mode, nullptr).Equal(lhs, rhs, map_free_vars);
});

TVM_REGISTER_GLOBAL("node.GetFirstStructuralMismatch")
.set_body_typed([](const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars) {
Optional<ObjectPathPair> first_mismatch;
bool equal = RemapVarSEqualHandler(false, &first_mismatch).Equal(lhs, rhs, map_free_vars);
bool equal = SEqualHandlerDefault(false, &first_mismatch).Equal(lhs, rhs, map_free_vars);
ICHECK(equal == !first_mismatch.defined());
return first_mismatch;
});

bool StructuralEqual::operator()(const ObjectRef& lhs, const ObjectRef& rhs) const {
return RemapVarSEqualHandler(false, nullptr).Equal(lhs, rhs, false);
return SEqualHandlerDefault(false, nullptr).Equal(lhs, rhs, false);
}

} // namespace tvm
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Force GitHub README to respect dark mode (function() { var style = document.createElement('style'); style.textContent = ' .markdown-body { color-scheme: dark light; } .markdown-body pre { background: #161b22 !important; } .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; } .markdown-body table th, .markdown-body table td { border-color: #30363d !important; } .markdown-body img { background: #0d1117; } .markdown-body blockquote { border-left-color: #8b949e; } .markdown-body hr { border-color: #30363d; } '; document.head.appendChild(style); })(); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 42 additions & 0 deletions include/tvm/node/structural_equal.h
Original file line numberDiff line numberDiff line change
Expand Up@@ -324,5 +324,47 @@ class SEqualReducer {
bool map_free_vars_ = false;
};

/*! \brief The default handler for equality testing.
*
* Users can derive from this class and override the DispatchSEqualReduce method,
* to customize equality testing.
*/
class SEqualHandlerDefault : public SEqualReducer::Handler {
public:
SEqualHandlerDefault(bool assert_mode, Optional<ObjectPathPair>* first_mismatch);
virtual ~SEqualHandlerDefault();

bool SEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) override;
void DeferFail(const ObjectPathPair& mismatch_paths) override;
ObjectRef MapLhsToRhs(const ObjectRef& lhs) override;
void MarkGraphNode() override;

/*!
* \brief The entry point for equality testing
* \param lhs The left operand.
* \param rhs The right operand.
* \param map_free_vars Whether or not to remap variables if possible.
* \return The equality result.
*/
virtual bool Equal(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars);

protected:
/*!
* \brief The dispatcher for equality testing of intermediate objects
* \param lhs The left operand.
* \param rhs The right operand.
* \param map_free_vars Whether or not to remap variables if possible.
* \param current_paths Optional paths to `lhs` and `rhs` objects, for error traceability.
* \return The equality result.
*/
virtual bool DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths);

private:
class Impl;
Impl* impl;
};

} // namespace tvm
#endif // TVM_NODE_STRUCTURAL_EQUAL_H_
38 changes: 38 additions & 0 deletions include/tvm/node/structural_hash.h
Original file line numberDiff line numberDiff line change
Expand Up@@ -200,6 +200,44 @@ class SHashReducer {
bool map_free_vars_;
};

/*! \brief The default handler for hash key computation
*
* Users can derive from this class and override the DispatchSHash method,
* to customize hashing.
*/
class SHashHandlerDefault : public SHashReducer::Handler {
public:
SHashHandlerDefault();
virtual ~SHashHandlerDefault();

void SHashReduceHashedValue(size_t hashed_value) override;
void SHashReduce(const ObjectRef& key, bool map_free_vars) override;
void SHashReduceFreeVar(const runtime::Object* var, bool map_free_vars) override;
bool LookupHashedValue(const ObjectRef& key, size_t* hashed_value) override;
void MarkGraphNode() override;

/*!
* \brief The entry point for hashing
* \param object The object to be hashed.
* \param map_free_vars Whether or not to remap variables if possible.
* \return The hash result.
*/
virtual size_t Hash(const ObjectRef& object, bool map_free_vars);

protected:
/*!
* \brief The dispatcher for hashing of intermediate objects
* \param object An intermediate object to be hashed.
* \param map_free_vars Whether or not to remap variables if possible.
* \return The hash result.
*/
virtual void DispatchSHash(const ObjectRef& object, bool map_free_vars);

private:
class Impl;
Impl* impl;
};

class SEqualReducer;
struct NDArrayContainerTrait {
static constexpr const std::nullptr_t VisitAttrs = nullptr;
Expand Down
109 changes: 73 additions & 36 deletions src/node/structural_equal.cc
Original file line numberDiff line numberDiff line change
Expand Up@@ -198,13 +198,13 @@ bool SEqualReducer::ObjectAttrsEqual(const ObjectRef& lhs, const ObjectRef& rhs,
* The order of SEqual being called is the same as the order as if we
* eagerly do recursive calls in SEqualReduce.
*/
class RemapVarSEqualHandler : public SEqualReducer::Handler {
class SEqualHandlerDefault::Impl {
public:
explicit RemapVarSEqualHandler(bool assert_mode, Optional<ObjectPathPair>* first_mismatch)
: assert_mode_(assert_mode), first_mismatch_(first_mismatch) {}
Impl(SEqualHandlerDefault* parent, bool assert_mode, Optional<ObjectPathPair>* first_mismatch)
: parent_(parent), assert_mode_(assert_mode), first_mismatch_(first_mismatch) {}

bool SEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) final {
const Optional<ObjectPathPair>& current_paths) {
// We cannot use check lhs.same_as(rhs) to check equality.
// if we choose to enable var remapping.
//
Expand DownExpand Up@@ -239,17 +239,17 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
return CheckResult(run(), lhs, rhs, current_paths);
}

void DeferFail(const ObjectPathPair& mismatch_paths) final {
void DeferFail(const ObjectPathPair& mismatch_paths) {
pending_tasks_.emplace_back(Task::ForceFailTag{}, mismatch_paths);
}

void MarkGraphNode() final {
void MarkGraphNode() {
// need to push to pending tasks in this case
ICHECK(!allow_push_to_stack_ && !task_stack_.empty());
task_stack_.back().graph_equal = true;
}

ObjectRef MapLhsToRhs(const ObjectRef& lhs) final {
ObjectRef MapLhsToRhs(const ObjectRef& lhs) {
auto it = equal_map_lhs_.find(lhs);
if (it != equal_map_lhs_.end()) return it->second;
return ObjectRef(nullptr);
Expand DownExpand Up@@ -279,7 +279,35 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
return RunTasks();
}

// The default equal as registered in the structural equal vtable.
bool DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
auto compute = [=]() {
ICHECK(lhs.defined() && rhs.defined() && lhs->type_index() == rhs->type_index());
// skip entries that already have equality maps.
auto it = equal_map_lhs_.find(lhs);
if (it != equal_map_lhs_.end()) {
return it->second.same_as(rhs);
}
if (equal_map_rhs_.count(rhs)) return false;

SEqualReducer reducer = GetReducer(lhs, rhs, map_free_vars, current_paths);
return vtable_->SEqualReduce(lhs.get(), rhs.get(), reducer);
};
return CheckResult(compute(), lhs, rhs, current_paths);
}

protected:
SEqualReducer GetReducer(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
if (!IsPathTracingEnabled()) {
return SEqualReducer(parent_, nullptr, map_free_vars);
} else {
PathTracingData tracing_data = {current_paths.value(), lhs, rhs, first_mismatch_};
return SEqualReducer(parent_, &tracing_data, map_free_vars);
}
}

// Check the result.
bool CheckResult(bool result, const ObjectRef& lhs, const ObjectRef& rhs,
const Optional<ObjectPathPair>& current_paths) {
Expand DownExpand Up@@ -335,7 +363,8 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
// which populates the pending tasks.
ICHECK_EQ(pending_tasks_.size(), 0U);
allow_push_to_stack_ = false;
if (!DispatchSEqualReduce(entry.lhs, entry.rhs, entry.map_free_vars, entry.current_paths))
if (!parent_->DispatchSEqualReduce(entry.lhs, entry.rhs, entry.map_free_vars,
entry.current_paths))
return false;
allow_push_to_stack_ = true;
// Push pending tasks in reverse order, so earlier tasks get to
Expand All@@ -349,31 +378,6 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
return true;
}

// The default equal as registered in the structural equal vtable.
bool DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
auto compute = [=]() {
ICHECK(lhs.defined() && rhs.defined() && lhs->type_index() == rhs->type_index());
// skip entries that already have equality maps.
auto it = equal_map_lhs_.find(lhs);
if (it != equal_map_lhs_.end()) {
return it->second.same_as(rhs);
}
if (equal_map_rhs_.count(rhs)) return false;

// Run reduce check for free nodes.
if (!IsPathTracingEnabled()) {
return vtable_->SEqualReduce(lhs.get(), rhs.get(),
SEqualReducer(this, nullptr, map_free_vars));
} else {
PathTracingData tracing_data = {current_paths.value(), lhs, rhs, first_mismatch_};
return vtable_->SEqualReduce(lhs.get(), rhs.get(),
SEqualReducer(this, &tracing_data, map_free_vars));
}
};
return CheckResult(compute(), lhs, rhs, current_paths);
}

private:
/*! \brief Pending reduce tasks. */
struct Task {
Expand DownExpand Up@@ -407,6 +411,8 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {

bool IsPathTracingEnabled() const { return first_mismatch_ != nullptr; }

// The owner of this impl
SEqualHandlerDefault* parent_;
// list of pending tasks to be pushed to the stack.
std::vector<Task> pending_tasks_;
// Internal task stack to executed the task.
Expand All@@ -425,22 +431,53 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
std::unordered_map<ObjectRef, ObjectRef, ObjectPtrHash, ObjectPtrEqual> equal_map_rhs_;
};

SEqualHandlerDefault::SEqualHandlerDefault(bool assert_mode,
Optional<ObjectPathPair>* first_mismatch) {
impl = new Impl(this, assert_mode, first_mismatch);
}

SEqualHandlerDefault::~SEqualHandlerDefault() { delete impl; }

bool SEqualHandlerDefault::SEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs,
bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
return impl->SEqualReduce(lhs, rhs, map_free_vars, current_paths);
}

void SEqualHandlerDefault::DeferFail(const ObjectPathPair& mismatch_paths) {
impl->DeferFail(mismatch_paths);
}

ObjectRef SEqualHandlerDefault::MapLhsToRhs(const ObjectRef& lhs) { return impl->MapLhsToRhs(lhs); }

void SEqualHandlerDefault::MarkGraphNode() { impl->MarkGraphNode(); }

bool SEqualHandlerDefault::Equal(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars) {
return impl->Equal(lhs, rhs, map_free_vars);
}

bool SEqualHandlerDefault::DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs,
bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
return impl->DispatchSEqualReduce(lhs, rhs, map_free_vars, current_paths);
}

TVM_REGISTER_GLOBAL("node.StructuralEqual")
.set_body_typed([](const ObjectRef& lhs, const ObjectRef& rhs, bool assert_mode,
bool map_free_vars) {
return RemapVarSEqualHandler(assert_mode, nullptr).Equal(lhs, rhs, map_free_vars);
return SEqualHandlerDefault(assert_mode, nullptr).Equal(lhs, rhs, map_free_vars);
});

TVM_REGISTER_GLOBAL("node.GetFirstStructuralMismatch")
.set_body_typed([](const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars) {
Optional<ObjectPathPair> first_mismatch;
bool equal = RemapVarSEqualHandler(false, &first_mismatch).Equal(lhs, rhs, map_free_vars);
bool equal = SEqualHandlerDefault(false, &first_mismatch).Equal(lhs, rhs, map_free_vars);
ICHECK(equal == !first_mismatch.defined());
return first_mismatch;
});

bool StructuralEqual::operator()(const ObjectRef& lhs, const ObjectRef& rhs) const {
return RemapVarSEqualHandler(false, nullptr).Equal(lhs, rhs, false);
return SEqualHandlerDefault(false, nullptr).Equal(lhs, rhs, false);
}

} // namespace tvm
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Highlight search terms from Google/DuckDuckGo/Bing referrer (function() { var ref = document.referrer; var terms = []; if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) { var url = new URL(ref); var q = url.searchParams.get('q') || url.searchParams.get('p'); if (q) { terms = q.split(/\s+/).filter(function(t) { return t.length > 2; }); } } if (terms.length === 0) return; var style = document.createElement('style'); style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }'; document.head.appendChild(style); function highlight(node) { if (node.nodeType === 3) { // text node var text = node.textContent; var found = false; terms.forEach(function(term) { var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\]\\]/g, '\\') + ')', 'gi'); if (regex.test(text)) { found = true; var frag = document.createDocumentFragment(); var parts = text.split(regex); parts.forEach(function(part, i) { if (i % 2 === 0) { frag.appendChild(document.createTextNode(part)); } else { var span = document.createElement('span'); span.className = 'userscript-highlight'; span.textContent = part; frag.appendChild(span); } }); node.parentNode.replaceChild(frag, node); } }); } else if (node.nodeType === 1 && node.childNodes) { // element var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT']; if (!skipTags.includes(node.tagName)) { Array.from(node.childNodes).forEach(highlight); } } } highlight(document.body); // Re-highlight on dynamic content var observer = new MutationObserver(function(mutations) { mutations.forEach(function(m) { m.addedNodes.forEach(function(node) { if (node.nodeType === 1 || node.nodeType === 3) highlight(node); }); }); }); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 42 additions & 0 deletions include/tvm/node/structural_equal.h
Original file line numberDiff line numberDiff line change
Expand Up@@ -324,5 +324,47 @@ class SEqualReducer {
bool map_free_vars_ = false;
};

/*! \brief The default handler for equality testing.
*
* Users can derive from this class and override the DispatchSEqualReduce method,
* to customize equality testing.
*/
class SEqualHandlerDefault : public SEqualReducer::Handler {
public:
SEqualHandlerDefault(bool assert_mode, Optional<ObjectPathPair>* first_mismatch);
virtual ~SEqualHandlerDefault();

bool SEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) override;
void DeferFail(const ObjectPathPair& mismatch_paths) override;
ObjectRef MapLhsToRhs(const ObjectRef& lhs) override;
void MarkGraphNode() override;

/*!
* \brief The entry point for equality testing
* \param lhs The left operand.
* \param rhs The right operand.
* \param map_free_vars Whether or not to remap variables if possible.
* \return The equality result.
*/
virtual bool Equal(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars);

protected:
/*!
* \brief The dispatcher for equality testing of intermediate objects
* \param lhs The left operand.
* \param rhs The right operand.
* \param map_free_vars Whether or not to remap variables if possible.
* \param current_paths Optional paths to `lhs` and `rhs` objects, for error traceability.
* \return The equality result.
*/
virtual bool DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths);

private:
class Impl;
Impl* impl;
};

} // namespace tvm
#endif // TVM_NODE_STRUCTURAL_EQUAL_H_
38 changes: 38 additions & 0 deletions include/tvm/node/structural_hash.h
Original file line numberDiff line numberDiff line change
Expand Up@@ -200,6 +200,44 @@ class SHashReducer {
bool map_free_vars_;
};

/*! \brief The default handler for hash key computation
*
* Users can derive from this class and override the DispatchSHash method,
* to customize hashing.
*/
class SHashHandlerDefault : public SHashReducer::Handler {
public:
SHashHandlerDefault();
virtual ~SHashHandlerDefault();

void SHashReduceHashedValue(size_t hashed_value) override;
void SHashReduce(const ObjectRef& key, bool map_free_vars) override;
void SHashReduceFreeVar(const runtime::Object* var, bool map_free_vars) override;
bool LookupHashedValue(const ObjectRef& key, size_t* hashed_value) override;
void MarkGraphNode() override;

/*!
* \brief The entry point for hashing
* \param object The object to be hashed.
* \param map_free_vars Whether or not to remap variables if possible.
* \return The hash result.
*/
virtual size_t Hash(const ObjectRef& object, bool map_free_vars);

protected:
/*!
* \brief The dispatcher for hashing of intermediate objects
* \param object An intermediate object to be hashed.
* \param map_free_vars Whether or not to remap variables if possible.
* \return The hash result.
*/
virtual void DispatchSHash(const ObjectRef& object, bool map_free_vars);

private:
class Impl;
Impl* impl;
};

class SEqualReducer;
struct NDArrayContainerTrait {
static constexpr const std::nullptr_t VisitAttrs = nullptr;
Expand Down
109 changes: 73 additions & 36 deletions src/node/structural_equal.cc
Original file line numberDiff line numberDiff line change
Expand Up@@ -198,13 +198,13 @@ bool SEqualReducer::ObjectAttrsEqual(const ObjectRef& lhs, const ObjectRef& rhs,
* The order of SEqual being called is the same as the order as if we
* eagerly do recursive calls in SEqualReduce.
*/
class RemapVarSEqualHandler : public SEqualReducer::Handler {
class SEqualHandlerDefault::Impl {
public:
explicit RemapVarSEqualHandler(bool assert_mode, Optional<ObjectPathPair>* first_mismatch)
: assert_mode_(assert_mode), first_mismatch_(first_mismatch) {}
Impl(SEqualHandlerDefault* parent, bool assert_mode, Optional<ObjectPathPair>* first_mismatch)
: parent_(parent), assert_mode_(assert_mode), first_mismatch_(first_mismatch) {}

bool SEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) final {
const Optional<ObjectPathPair>& current_paths) {
// We cannot use check lhs.same_as(rhs) to check equality.
// if we choose to enable var remapping.
//
Expand DownExpand Up@@ -239,17 +239,17 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
return CheckResult(run(), lhs, rhs, current_paths);
}

void DeferFail(const ObjectPathPair& mismatch_paths) final {
void DeferFail(const ObjectPathPair& mismatch_paths) {
pending_tasks_.emplace_back(Task::ForceFailTag{}, mismatch_paths);
}

void MarkGraphNode() final {
void MarkGraphNode() {
// need to push to pending tasks in this case
ICHECK(!allow_push_to_stack_ && !task_stack_.empty());
task_stack_.back().graph_equal = true;
}

ObjectRef MapLhsToRhs(const ObjectRef& lhs) final {
ObjectRef MapLhsToRhs(const ObjectRef& lhs) {
auto it = equal_map_lhs_.find(lhs);
if (it != equal_map_lhs_.end()) return it->second;
return ObjectRef(nullptr);
Expand DownExpand Up@@ -279,7 +279,35 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
return RunTasks();
}

// The default equal as registered in the structural equal vtable.
bool DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
auto compute = [=]() {
ICHECK(lhs.defined() && rhs.defined() && lhs->type_index() == rhs->type_index());
// skip entries that already have equality maps.
auto it = equal_map_lhs_.find(lhs);
if (it != equal_map_lhs_.end()) {
return it->second.same_as(rhs);
}
if (equal_map_rhs_.count(rhs)) return false;

SEqualReducer reducer = GetReducer(lhs, rhs, map_free_vars, current_paths);
return vtable_->SEqualReduce(lhs.get(), rhs.get(), reducer);
};
return CheckResult(compute(), lhs, rhs, current_paths);
}

protected:
SEqualReducer GetReducer(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
if (!IsPathTracingEnabled()) {
return SEqualReducer(parent_, nullptr, map_free_vars);
} else {
PathTracingData tracing_data = {current_paths.value(), lhs, rhs, first_mismatch_};
return SEqualReducer(parent_, &tracing_data, map_free_vars);
}
}

// Check the result.
bool CheckResult(bool result, const ObjectRef& lhs, const ObjectRef& rhs,
const Optional<ObjectPathPair>& current_paths) {
Expand DownExpand Up@@ -335,7 +363,8 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
// which populates the pending tasks.
ICHECK_EQ(pending_tasks_.size(), 0U);
allow_push_to_stack_ = false;
if (!DispatchSEqualReduce(entry.lhs, entry.rhs, entry.map_free_vars, entry.current_paths))
if (!parent_->DispatchSEqualReduce(entry.lhs, entry.rhs, entry.map_free_vars,
entry.current_paths))
return false;
allow_push_to_stack_ = true;
// Push pending tasks in reverse order, so earlier tasks get to
Expand All@@ -349,31 +378,6 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
return true;
}

// The default equal as registered in the structural equal vtable.
bool DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
auto compute = [=]() {
ICHECK(lhs.defined() && rhs.defined() && lhs->type_index() == rhs->type_index());
// skip entries that already have equality maps.
auto it = equal_map_lhs_.find(lhs);
if (it != equal_map_lhs_.end()) {
return it->second.same_as(rhs);
}
if (equal_map_rhs_.count(rhs)) return false;

// Run reduce check for free nodes.
if (!IsPathTracingEnabled()) {
return vtable_->SEqualReduce(lhs.get(), rhs.get(),
SEqualReducer(this, nullptr, map_free_vars));
} else {
PathTracingData tracing_data = {current_paths.value(), lhs, rhs, first_mismatch_};
return vtable_->SEqualReduce(lhs.get(), rhs.get(),
SEqualReducer(this, &tracing_data, map_free_vars));
}
};
return CheckResult(compute(), lhs, rhs, current_paths);
}

private:
/*! \brief Pending reduce tasks. */
struct Task {
Expand DownExpand Up@@ -407,6 +411,8 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {

bool IsPathTracingEnabled() const { return first_mismatch_ != nullptr; }

// The owner of this impl
SEqualHandlerDefault* parent_;
// list of pending tasks to be pushed to the stack.
std::vector<Task> pending_tasks_;
// Internal task stack to executed the task.
Expand All@@ -425,22 +431,53 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
std::unordered_map<ObjectRef, ObjectRef, ObjectPtrHash, ObjectPtrEqual> equal_map_rhs_;
};

SEqualHandlerDefault::SEqualHandlerDefault(bool assert_mode,
Optional<ObjectPathPair>* first_mismatch) {
impl = new Impl(this, assert_mode, first_mismatch);
}

SEqualHandlerDefault::~SEqualHandlerDefault() { delete impl; }

bool SEqualHandlerDefault::SEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs,
bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
return impl->SEqualReduce(lhs, rhs, map_free_vars, current_paths);
}

void SEqualHandlerDefault::DeferFail(const ObjectPathPair& mismatch_paths) {
impl->DeferFail(mismatch_paths);
}

ObjectRef SEqualHandlerDefault::MapLhsToRhs(const ObjectRef& lhs) { return impl->MapLhsToRhs(lhs); }

void SEqualHandlerDefault::MarkGraphNode() { impl->MarkGraphNode(); }

bool SEqualHandlerDefault::Equal(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars) {
return impl->Equal(lhs, rhs, map_free_vars);
}

bool SEqualHandlerDefault::DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs,
bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
return impl->DispatchSEqualReduce(lhs, rhs, map_free_vars, current_paths);
}

TVM_REGISTER_GLOBAL("node.StructuralEqual")
.set_body_typed([](const ObjectRef& lhs, const ObjectRef& rhs, bool assert_mode,
bool map_free_vars) {
return RemapVarSEqualHandler(assert_mode, nullptr).Equal(lhs, rhs, map_free_vars);
return SEqualHandlerDefault(assert_mode, nullptr).Equal(lhs, rhs, map_free_vars);
});

TVM_REGISTER_GLOBAL("node.GetFirstStructuralMismatch")
.set_body_typed([](const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars) {
Optional<ObjectPathPair> first_mismatch;
bool equal = RemapVarSEqualHandler(false, &first_mismatch).Equal(lhs, rhs, map_free_vars);
bool equal = SEqualHandlerDefault(false, &first_mismatch).Equal(lhs, rhs, map_free_vars);
ICHECK(equal == !first_mismatch.defined());
return first_mismatch;
});

bool StructuralEqual::operator()(const ObjectRef& lhs, const ObjectRef& rhs) const {
return RemapVarSEqualHandler(false, nullptr).Equal(lhs, rhs, false);
return SEqualHandlerDefault(false, nullptr).Equal(lhs, rhs, false);
}

} // namespace tvm
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Strip utm_, fbclid, gclid, etc. from all links on page (function() { var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content', 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid', 'ref', 'ref_src', 'source', 'medium', 'campaign']; function cleanUrl(url) { try { var u = new URL(url, window.location.origin); var changed = false; trackingParams.forEach(function(p) { if (u.searchParams.has(p)) { u.searchParams.delete(p); changed = true; } }); return changed ? u.toString() : url; } catch (e) { return url; } } function cleanLinks() { document.querySelectorAll('a[href]').forEach(function(a) { var clean = cleanUrl(a.href); if (clean !== a.href) a.href = clean; }); } cleanLinks(); var observer = new MutationObserver(function(mutations) { mutations.forEach(function(m) { m.addedNodes.forEach(function(node) { if (node.nodeType === 1) { if (node.tagName === 'A') cleanLinks(); node.querySelectorAll('a[href]').forEach(function(a) { var clean = cleanUrl(a.href); if (clean !== a.href) a.href = clean; }); } }); }); }); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 42 additions & 0 deletions include/tvm/node/structural_equal.h
Original file line numberDiff line numberDiff line change
Expand Up@@ -324,5 +324,47 @@ class SEqualReducer {
bool map_free_vars_ = false;
};

/*! \brief The default handler for equality testing.
*
* Users can derive from this class and override the DispatchSEqualReduce method,
* to customize equality testing.
*/
class SEqualHandlerDefault : public SEqualReducer::Handler {
public:
SEqualHandlerDefault(bool assert_mode, Optional<ObjectPathPair>* first_mismatch);
virtual ~SEqualHandlerDefault();

bool SEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) override;
void DeferFail(const ObjectPathPair& mismatch_paths) override;
ObjectRef MapLhsToRhs(const ObjectRef& lhs) override;
void MarkGraphNode() override;

/*!
* \brief The entry point for equality testing
* \param lhs The left operand.
* \param rhs The right operand.
* \param map_free_vars Whether or not to remap variables if possible.
* \return The equality result.
*/
virtual bool Equal(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars);

protected:
/*!
* \brief The dispatcher for equality testing of intermediate objects
* \param lhs The left operand.
* \param rhs The right operand.
* \param map_free_vars Whether or not to remap variables if possible.
* \param current_paths Optional paths to `lhs` and `rhs` objects, for error traceability.
* \return The equality result.
*/
virtual bool DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths);

private:
class Impl;
Impl* impl;
};

} // namespace tvm
#endif // TVM_NODE_STRUCTURAL_EQUAL_H_
38 changes: 38 additions & 0 deletions include/tvm/node/structural_hash.h
Original file line numberDiff line numberDiff line change
Expand Up@@ -200,6 +200,44 @@ class SHashReducer {
bool map_free_vars_;
};

/*! \brief The default handler for hash key computation
*
* Users can derive from this class and override the DispatchSHash method,
* to customize hashing.
*/
class SHashHandlerDefault : public SHashReducer::Handler {
public:
SHashHandlerDefault();
virtual ~SHashHandlerDefault();

void SHashReduceHashedValue(size_t hashed_value) override;
void SHashReduce(const ObjectRef& key, bool map_free_vars) override;
void SHashReduceFreeVar(const runtime::Object* var, bool map_free_vars) override;
bool LookupHashedValue(const ObjectRef& key, size_t* hashed_value) override;
void MarkGraphNode() override;

/*!
* \brief The entry point for hashing
* \param object The object to be hashed.
* \param map_free_vars Whether or not to remap variables if possible.
* \return The hash result.
*/
virtual size_t Hash(const ObjectRef& object, bool map_free_vars);

protected:
/*!
* \brief The dispatcher for hashing of intermediate objects
* \param object An intermediate object to be hashed.
* \param map_free_vars Whether or not to remap variables if possible.
* \return The hash result.
*/
virtual void DispatchSHash(const ObjectRef& object, bool map_free_vars);

private:
class Impl;
Impl* impl;
};

class SEqualReducer;
struct NDArrayContainerTrait {
static constexpr const std::nullptr_t VisitAttrs = nullptr;
Expand Down
109 changes: 73 additions & 36 deletions src/node/structural_equal.cc
Original file line numberDiff line numberDiff line change
Expand Up@@ -198,13 +198,13 @@ bool SEqualReducer::ObjectAttrsEqual(const ObjectRef& lhs, const ObjectRef& rhs,
* The order of SEqual being called is the same as the order as if we
* eagerly do recursive calls in SEqualReduce.
*/
class RemapVarSEqualHandler : public SEqualReducer::Handler {
class SEqualHandlerDefault::Impl {
public:
explicit RemapVarSEqualHandler(bool assert_mode, Optional<ObjectPathPair>* first_mismatch)
: assert_mode_(assert_mode), first_mismatch_(first_mismatch) {}
Impl(SEqualHandlerDefault* parent, bool assert_mode, Optional<ObjectPathPair>* first_mismatch)
: parent_(parent), assert_mode_(assert_mode), first_mismatch_(first_mismatch) {}

bool SEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) final {
const Optional<ObjectPathPair>& current_paths) {
// We cannot use check lhs.same_as(rhs) to check equality.
// if we choose to enable var remapping.
//
Expand DownExpand Up@@ -239,17 +239,17 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
return CheckResult(run(), lhs, rhs, current_paths);
}

void DeferFail(const ObjectPathPair& mismatch_paths) final {
void DeferFail(const ObjectPathPair& mismatch_paths) {
pending_tasks_.emplace_back(Task::ForceFailTag{}, mismatch_paths);
}

void MarkGraphNode() final {
void MarkGraphNode() {
// need to push to pending tasks in this case
ICHECK(!allow_push_to_stack_ && !task_stack_.empty());
task_stack_.back().graph_equal = true;
}

ObjectRef MapLhsToRhs(const ObjectRef& lhs) final {
ObjectRef MapLhsToRhs(const ObjectRef& lhs) {
auto it = equal_map_lhs_.find(lhs);
if (it != equal_map_lhs_.end()) return it->second;
return ObjectRef(nullptr);
Expand DownExpand Up@@ -279,7 +279,35 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
return RunTasks();
}

// The default equal as registered in the structural equal vtable.
bool DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
auto compute = [=]() {
ICHECK(lhs.defined() && rhs.defined() && lhs->type_index() == rhs->type_index());
// skip entries that already have equality maps.
auto it = equal_map_lhs_.find(lhs);
if (it != equal_map_lhs_.end()) {
return it->second.same_as(rhs);
}
if (equal_map_rhs_.count(rhs)) return false;

SEqualReducer reducer = GetReducer(lhs, rhs, map_free_vars, current_paths);
return vtable_->SEqualReduce(lhs.get(), rhs.get(), reducer);
};
return CheckResult(compute(), lhs, rhs, current_paths);
}

protected:
SEqualReducer GetReducer(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
if (!IsPathTracingEnabled()) {
return SEqualReducer(parent_, nullptr, map_free_vars);
} else {
PathTracingData tracing_data = {current_paths.value(), lhs, rhs, first_mismatch_};
return SEqualReducer(parent_, &tracing_data, map_free_vars);
}
}

// Check the result.
bool CheckResult(bool result, const ObjectRef& lhs, const ObjectRef& rhs,
const Optional<ObjectPathPair>& current_paths) {
Expand DownExpand Up@@ -335,7 +363,8 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
// which populates the pending tasks.
ICHECK_EQ(pending_tasks_.size(), 0U);
allow_push_to_stack_ = false;
if (!DispatchSEqualReduce(entry.lhs, entry.rhs, entry.map_free_vars, entry.current_paths))
if (!parent_->DispatchSEqualReduce(entry.lhs, entry.rhs, entry.map_free_vars,
entry.current_paths))
return false;
allow_push_to_stack_ = true;
// Push pending tasks in reverse order, so earlier tasks get to
Expand All@@ -349,31 +378,6 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
return true;
}

// The default equal as registered in the structural equal vtable.
bool DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
auto compute = [=]() {
ICHECK(lhs.defined() && rhs.defined() && lhs->type_index() == rhs->type_index());
// skip entries that already have equality maps.
auto it = equal_map_lhs_.find(lhs);
if (it != equal_map_lhs_.end()) {
return it->second.same_as(rhs);
}
if (equal_map_rhs_.count(rhs)) return false;

// Run reduce check for free nodes.
if (!IsPathTracingEnabled()) {
return vtable_->SEqualReduce(lhs.get(), rhs.get(),
SEqualReducer(this, nullptr, map_free_vars));
} else {
PathTracingData tracing_data = {current_paths.value(), lhs, rhs, first_mismatch_};
return vtable_->SEqualReduce(lhs.get(), rhs.get(),
SEqualReducer(this, &tracing_data, map_free_vars));
}
};
return CheckResult(compute(), lhs, rhs, current_paths);
}

private:
/*! \brief Pending reduce tasks. */
struct Task {
Expand DownExpand Up@@ -407,6 +411,8 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {

bool IsPathTracingEnabled() const { return first_mismatch_ != nullptr; }

// The owner of this impl
SEqualHandlerDefault* parent_;
// list of pending tasks to be pushed to the stack.
std::vector<Task> pending_tasks_;
// Internal task stack to executed the task.
Expand All@@ -425,22 +431,53 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
std::unordered_map<ObjectRef, ObjectRef, ObjectPtrHash, ObjectPtrEqual> equal_map_rhs_;
};

SEqualHandlerDefault::SEqualHandlerDefault(bool assert_mode,
Optional<ObjectPathPair>* first_mismatch) {
impl = new Impl(this, assert_mode, first_mismatch);
}

SEqualHandlerDefault::~SEqualHandlerDefault() { delete impl; }

bool SEqualHandlerDefault::SEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs,
bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
return impl->SEqualReduce(lhs, rhs, map_free_vars, current_paths);
}

void SEqualHandlerDefault::DeferFail(const ObjectPathPair& mismatch_paths) {
impl->DeferFail(mismatch_paths);
}

ObjectRef SEqualHandlerDefault::MapLhsToRhs(const ObjectRef& lhs) { return impl->MapLhsToRhs(lhs); }

void SEqualHandlerDefault::MarkGraphNode() { impl->MarkGraphNode(); }

bool SEqualHandlerDefault::Equal(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars) {
return impl->Equal(lhs, rhs, map_free_vars);
}

bool SEqualHandlerDefault::DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs,
bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
return impl->DispatchSEqualReduce(lhs, rhs, map_free_vars, current_paths);
}

TVM_REGISTER_GLOBAL("node.StructuralEqual")
.set_body_typed([](const ObjectRef& lhs, const ObjectRef& rhs, bool assert_mode,
bool map_free_vars) {
return RemapVarSEqualHandler(assert_mode, nullptr).Equal(lhs, rhs, map_free_vars);
return SEqualHandlerDefault(assert_mode, nullptr).Equal(lhs, rhs, map_free_vars);
});

TVM_REGISTER_GLOBAL("node.GetFirstStructuralMismatch")
.set_body_typed([](const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars) {
Optional<ObjectPathPair> first_mismatch;
bool equal = RemapVarSEqualHandler(false, &first_mismatch).Equal(lhs, rhs, map_free_vars);
bool equal = SEqualHandlerDefault(false, &first_mismatch).Equal(lhs, rhs, map_free_vars);
ICHECK(equal == !first_mismatch.defined());
return first_mismatch;
});

bool StructuralEqual::operator()(const ObjectRef& lhs, const ObjectRef& rhs) const {
return RemapVarSEqualHandler(false, nullptr).Equal(lhs, rhs, false);
return SEqualHandlerDefault(false, nullptr).Equal(lhs, rhs, false);
}

} // namespace tvm
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Auto-enable theater mode on YouTube (function() { function tryTheater() { var btn = document.querySelector('button[aria-label="Theater mode"], ytd-player #player button[title="Theater mode"]'); if (btn && !btn.classList.contains('activated')) { btn.click(); } } // Try immediately tryTheater(); // Try after navigation (SPA) var lastUrl = location.href; setInterval(function() { if (location.href !== lastUrl) { lastUrl = location.href; setTimeout(tryTheater, 500); } }, 1000); // Also try on player load var observer = new MutationObserver(tryTheater); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 42 additions & 0 deletions include/tvm/node/structural_equal.h
Original file line numberDiff line numberDiff line change
Expand Up@@ -324,5 +324,47 @@ class SEqualReducer {
bool map_free_vars_ = false;
};

/*! \brief The default handler for equality testing.
*
* Users can derive from this class and override the DispatchSEqualReduce method,
* to customize equality testing.
*/
class SEqualHandlerDefault : public SEqualReducer::Handler {
public:
SEqualHandlerDefault(bool assert_mode, Optional<ObjectPathPair>* first_mismatch);
virtual ~SEqualHandlerDefault();

bool SEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) override;
void DeferFail(const ObjectPathPair& mismatch_paths) override;
ObjectRef MapLhsToRhs(const ObjectRef& lhs) override;
void MarkGraphNode() override;

/*!
* \brief The entry point for equality testing
* \param lhs The left operand.
* \param rhs The right operand.
* \param map_free_vars Whether or not to remap variables if possible.
* \return The equality result.
*/
virtual bool Equal(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars);

protected:
/*!
* \brief The dispatcher for equality testing of intermediate objects
* \param lhs The left operand.
* \param rhs The right operand.
* \param map_free_vars Whether or not to remap variables if possible.
* \param current_paths Optional paths to `lhs` and `rhs` objects, for error traceability.
* \return The equality result.
*/
virtual bool DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths);

private:
class Impl;
Impl* impl;
};

} // namespace tvm
#endif // TVM_NODE_STRUCTURAL_EQUAL_H_
38 changes: 38 additions & 0 deletions include/tvm/node/structural_hash.h
Original file line numberDiff line numberDiff line change
Expand Up@@ -200,6 +200,44 @@ class SHashReducer {
bool map_free_vars_;
};

/*! \brief The default handler for hash key computation
*
* Users can derive from this class and override the DispatchSHash method,
* to customize hashing.
*/
class SHashHandlerDefault : public SHashReducer::Handler {
public:
SHashHandlerDefault();
virtual ~SHashHandlerDefault();

void SHashReduceHashedValue(size_t hashed_value) override;
void SHashReduce(const ObjectRef& key, bool map_free_vars) override;
void SHashReduceFreeVar(const runtime::Object* var, bool map_free_vars) override;
bool LookupHashedValue(const ObjectRef& key, size_t* hashed_value) override;
void MarkGraphNode() override;

/*!
* \brief The entry point for hashing
* \param object The object to be hashed.
* \param map_free_vars Whether or not to remap variables if possible.
* \return The hash result.
*/
virtual size_t Hash(const ObjectRef& object, bool map_free_vars);

protected:
/*!
* \brief The dispatcher for hashing of intermediate objects
* \param object An intermediate object to be hashed.
* \param map_free_vars Whether or not to remap variables if possible.
* \return The hash result.
*/
virtual void DispatchSHash(const ObjectRef& object, bool map_free_vars);

private:
class Impl;
Impl* impl;
};

class SEqualReducer;
struct NDArrayContainerTrait {
static constexpr const std::nullptr_t VisitAttrs = nullptr;
Expand Down
109 changes: 73 additions & 36 deletions src/node/structural_equal.cc
Original file line numberDiff line numberDiff line change
Expand Up@@ -198,13 +198,13 @@ bool SEqualReducer::ObjectAttrsEqual(const ObjectRef& lhs, const ObjectRef& rhs,
* The order of SEqual being called is the same as the order as if we
* eagerly do recursive calls in SEqualReduce.
*/
class RemapVarSEqualHandler : public SEqualReducer::Handler {
class SEqualHandlerDefault::Impl {
public:
explicit RemapVarSEqualHandler(bool assert_mode, Optional<ObjectPathPair>* first_mismatch)
: assert_mode_(assert_mode), first_mismatch_(first_mismatch) {}
Impl(SEqualHandlerDefault* parent, bool assert_mode, Optional<ObjectPathPair>* first_mismatch)
: parent_(parent), assert_mode_(assert_mode), first_mismatch_(first_mismatch) {}

bool SEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) final {
const Optional<ObjectPathPair>& current_paths) {
// We cannot use check lhs.same_as(rhs) to check equality.
// if we choose to enable var remapping.
//
Expand DownExpand Up@@ -239,17 +239,17 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
return CheckResult(run(), lhs, rhs, current_paths);
}

void DeferFail(const ObjectPathPair& mismatch_paths) final {
void DeferFail(const ObjectPathPair& mismatch_paths) {
pending_tasks_.emplace_back(Task::ForceFailTag{}, mismatch_paths);
}

void MarkGraphNode() final {
void MarkGraphNode() {
// need to push to pending tasks in this case
ICHECK(!allow_push_to_stack_ && !task_stack_.empty());
task_stack_.back().graph_equal = true;
}

ObjectRef MapLhsToRhs(const ObjectRef& lhs) final {
ObjectRef MapLhsToRhs(const ObjectRef& lhs) {
auto it = equal_map_lhs_.find(lhs);
if (it != equal_map_lhs_.end()) return it->second;
return ObjectRef(nullptr);
Expand DownExpand Up@@ -279,7 +279,35 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
return RunTasks();
}

// The default equal as registered in the structural equal vtable.
bool DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
auto compute = [=]() {
ICHECK(lhs.defined() && rhs.defined() && lhs->type_index() == rhs->type_index());
// skip entries that already have equality maps.
auto it = equal_map_lhs_.find(lhs);
if (it != equal_map_lhs_.end()) {
return it->second.same_as(rhs);
}
if (equal_map_rhs_.count(rhs)) return false;

SEqualReducer reducer = GetReducer(lhs, rhs, map_free_vars, current_paths);
return vtable_->SEqualReduce(lhs.get(), rhs.get(), reducer);
};
return CheckResult(compute(), lhs, rhs, current_paths);
}

protected:
SEqualReducer GetReducer(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
if (!IsPathTracingEnabled()) {
return SEqualReducer(parent_, nullptr, map_free_vars);
} else {
PathTracingData tracing_data = {current_paths.value(), lhs, rhs, first_mismatch_};
return SEqualReducer(parent_, &tracing_data, map_free_vars);
}
}

// Check the result.
bool CheckResult(bool result, const ObjectRef& lhs, const ObjectRef& rhs,
const Optional<ObjectPathPair>& current_paths) {
Expand DownExpand Up@@ -335,7 +363,8 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
// which populates the pending tasks.
ICHECK_EQ(pending_tasks_.size(), 0U);
allow_push_to_stack_ = false;
if (!DispatchSEqualReduce(entry.lhs, entry.rhs, entry.map_free_vars, entry.current_paths))
if (!parent_->DispatchSEqualReduce(entry.lhs, entry.rhs, entry.map_free_vars,
entry.current_paths))
return false;
allow_push_to_stack_ = true;
// Push pending tasks in reverse order, so earlier tasks get to
Expand All@@ -349,31 +378,6 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
return true;
}

// The default equal as registered in the structural equal vtable.
bool DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
auto compute = [=]() {
ICHECK(lhs.defined() && rhs.defined() && lhs->type_index() == rhs->type_index());
// skip entries that already have equality maps.
auto it = equal_map_lhs_.find(lhs);
if (it != equal_map_lhs_.end()) {
return it->second.same_as(rhs);
}
if (equal_map_rhs_.count(rhs)) return false;

// Run reduce check for free nodes.
if (!IsPathTracingEnabled()) {
return vtable_->SEqualReduce(lhs.get(), rhs.get(),
SEqualReducer(this, nullptr, map_free_vars));
} else {
PathTracingData tracing_data = {current_paths.value(), lhs, rhs, first_mismatch_};
return vtable_->SEqualReduce(lhs.get(), rhs.get(),
SEqualReducer(this, &tracing_data, map_free_vars));
}
};
return CheckResult(compute(), lhs, rhs, current_paths);
}

private:
/*! \brief Pending reduce tasks. */
struct Task {
Expand DownExpand Up@@ -407,6 +411,8 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {

bool IsPathTracingEnabled() const { return first_mismatch_ != nullptr; }

// The owner of this impl
SEqualHandlerDefault* parent_;
// list of pending tasks to be pushed to the stack.
std::vector<Task> pending_tasks_;
// Internal task stack to executed the task.
Expand All@@ -425,22 +431,53 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
std::unordered_map<ObjectRef, ObjectRef, ObjectPtrHash, ObjectPtrEqual> equal_map_rhs_;
};

SEqualHandlerDefault::SEqualHandlerDefault(bool assert_mode,
Optional<ObjectPathPair>* first_mismatch) {
impl = new Impl(this, assert_mode, first_mismatch);
}

SEqualHandlerDefault::~SEqualHandlerDefault() { delete impl; }

bool SEqualHandlerDefault::SEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs,
bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
return impl->SEqualReduce(lhs, rhs, map_free_vars, current_paths);
}

void SEqualHandlerDefault::DeferFail(const ObjectPathPair& mismatch_paths) {
impl->DeferFail(mismatch_paths);
}

ObjectRef SEqualHandlerDefault::MapLhsToRhs(const ObjectRef& lhs) { return impl->MapLhsToRhs(lhs); }

void SEqualHandlerDefault::MarkGraphNode() { impl->MarkGraphNode(); }

bool SEqualHandlerDefault::Equal(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars) {
return impl->Equal(lhs, rhs, map_free_vars);
}

bool SEqualHandlerDefault::DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs,
bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
return impl->DispatchSEqualReduce(lhs, rhs, map_free_vars, current_paths);
}

TVM_REGISTER_GLOBAL("node.StructuralEqual")
.set_body_typed([](const ObjectRef& lhs, const ObjectRef& rhs, bool assert_mode,
bool map_free_vars) {
return RemapVarSEqualHandler(assert_mode, nullptr).Equal(lhs, rhs, map_free_vars);
return SEqualHandlerDefault(assert_mode, nullptr).Equal(lhs, rhs, map_free_vars);
});

TVM_REGISTER_GLOBAL("node.GetFirstStructuralMismatch")
.set_body_typed([](const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars) {
Optional<ObjectPathPair> first_mismatch;
bool equal = RemapVarSEqualHandler(false, &first_mismatch).Equal(lhs, rhs, map_free_vars);
bool equal = SEqualHandlerDefault(false, &first_mismatch).Equal(lhs, rhs, map_free_vars);
ICHECK(equal == !first_mismatch.defined());
return first_mismatch;
});

bool StructuralEqual::operator()(const ObjectRef& lhs, const ObjectRef& rhs) const {
return RemapVarSEqualHandler(false, nullptr).Equal(lhs, rhs, false);
return SEqualHandlerDefault(false, nullptr).Equal(lhs, rhs, false);
}

} // namespace tvm
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Remove or un-stick sticky/fixed headers that block content (function() { function unstick() { document.querySelectorAll('header, nav, [role="banner"], .header, .navbar, .sticky, .fixed-top, [style*="position: fixed"], [style*="position:sticky"]').forEach(function(el) { if (el.style.position === 'fixed' || el.style.position === 'sticky' || getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') { el.style.position = 'static'; el.style.top = 'auto'; el.style.zIndex = 'auto'; } }); } unstick(); var observer = new MutationObserver(unstick); observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] }); })(); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 42 additions & 0 deletions include/tvm/node/structural_equal.h
Original file line numberDiff line numberDiff line change
Expand Up@@ -324,5 +324,47 @@ class SEqualReducer {
bool map_free_vars_ = false;
};

/*! \brief The default handler for equality testing.
*
* Users can derive from this class and override the DispatchSEqualReduce method,
* to customize equality testing.
*/
class SEqualHandlerDefault : public SEqualReducer::Handler {
public:
SEqualHandlerDefault(bool assert_mode, Optional<ObjectPathPair>* first_mismatch);
virtual ~SEqualHandlerDefault();

bool SEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) override;
void DeferFail(const ObjectPathPair& mismatch_paths) override;
ObjectRef MapLhsToRhs(const ObjectRef& lhs) override;
void MarkGraphNode() override;

/*!
* \brief The entry point for equality testing
* \param lhs The left operand.
* \param rhs The right operand.
* \param map_free_vars Whether or not to remap variables if possible.
* \return The equality result.
*/
virtual bool Equal(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars);

protected:
/*!
* \brief The dispatcher for equality testing of intermediate objects
* \param lhs The left operand.
* \param rhs The right operand.
* \param map_free_vars Whether or not to remap variables if possible.
* \param current_paths Optional paths to `lhs` and `rhs` objects, for error traceability.
* \return The equality result.
*/
virtual bool DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths);

private:
class Impl;
Impl* impl;
};

} // namespace tvm
#endif // TVM_NODE_STRUCTURAL_EQUAL_H_
38 changes: 38 additions & 0 deletions include/tvm/node/structural_hash.h
Original file line numberDiff line numberDiff line change
Expand Up@@ -200,6 +200,44 @@ class SHashReducer {
bool map_free_vars_;
};

/*! \brief The default handler for hash key computation
*
* Users can derive from this class and override the DispatchSHash method,
* to customize hashing.
*/
class SHashHandlerDefault : public SHashReducer::Handler {
public:
SHashHandlerDefault();
virtual ~SHashHandlerDefault();

void SHashReduceHashedValue(size_t hashed_value) override;
void SHashReduce(const ObjectRef& key, bool map_free_vars) override;
void SHashReduceFreeVar(const runtime::Object* var, bool map_free_vars) override;
bool LookupHashedValue(const ObjectRef& key, size_t* hashed_value) override;
void MarkGraphNode() override;

/*!
* \brief The entry point for hashing
* \param object The object to be hashed.
* \param map_free_vars Whether or not to remap variables if possible.
* \return The hash result.
*/
virtual size_t Hash(const ObjectRef& object, bool map_free_vars);

protected:
/*!
* \brief The dispatcher for hashing of intermediate objects
* \param object An intermediate object to be hashed.
* \param map_free_vars Whether or not to remap variables if possible.
* \return The hash result.
*/
virtual void DispatchSHash(const ObjectRef& object, bool map_free_vars);

private:
class Impl;
Impl* impl;
};

class SEqualReducer;
struct NDArrayContainerTrait {
static constexpr const std::nullptr_t VisitAttrs = nullptr;
Expand Down
109 changes: 73 additions & 36 deletions src/node/structural_equal.cc
Original file line numberDiff line numberDiff line change
Expand Up@@ -198,13 +198,13 @@ bool SEqualReducer::ObjectAttrsEqual(const ObjectRef& lhs, const ObjectRef& rhs,
* The order of SEqual being called is the same as the order as if we
* eagerly do recursive calls in SEqualReduce.
*/
class RemapVarSEqualHandler : public SEqualReducer::Handler {
class SEqualHandlerDefault::Impl {
public:
explicit RemapVarSEqualHandler(bool assert_mode, Optional<ObjectPathPair>* first_mismatch)
: assert_mode_(assert_mode), first_mismatch_(first_mismatch) {}
Impl(SEqualHandlerDefault* parent, bool assert_mode, Optional<ObjectPathPair>* first_mismatch)
: parent_(parent), assert_mode_(assert_mode), first_mismatch_(first_mismatch) {}

bool SEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) final {
const Optional<ObjectPathPair>& current_paths) {
// We cannot use check lhs.same_as(rhs) to check equality.
// if we choose to enable var remapping.
//
Expand DownExpand Up@@ -239,17 +239,17 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
return CheckResult(run(), lhs, rhs, current_paths);
}

void DeferFail(const ObjectPathPair& mismatch_paths) final {
void DeferFail(const ObjectPathPair& mismatch_paths) {
pending_tasks_.emplace_back(Task::ForceFailTag{}, mismatch_paths);
}

void MarkGraphNode() final {
void MarkGraphNode() {
// need to push to pending tasks in this case
ICHECK(!allow_push_to_stack_ && !task_stack_.empty());
task_stack_.back().graph_equal = true;
}

ObjectRef MapLhsToRhs(const ObjectRef& lhs) final {
ObjectRef MapLhsToRhs(const ObjectRef& lhs) {
auto it = equal_map_lhs_.find(lhs);
if (it != equal_map_lhs_.end()) return it->second;
return ObjectRef(nullptr);
Expand DownExpand Up@@ -279,7 +279,35 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
return RunTasks();
}

// The default equal as registered in the structural equal vtable.
bool DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
auto compute = [=]() {
ICHECK(lhs.defined() && rhs.defined() && lhs->type_index() == rhs->type_index());
// skip entries that already have equality maps.
auto it = equal_map_lhs_.find(lhs);
if (it != equal_map_lhs_.end()) {
return it->second.same_as(rhs);
}
if (equal_map_rhs_.count(rhs)) return false;

SEqualReducer reducer = GetReducer(lhs, rhs, map_free_vars, current_paths);
return vtable_->SEqualReduce(lhs.get(), rhs.get(), reducer);
};
return CheckResult(compute(), lhs, rhs, current_paths);
}

protected:
SEqualReducer GetReducer(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
if (!IsPathTracingEnabled()) {
return SEqualReducer(parent_, nullptr, map_free_vars);
} else {
PathTracingData tracing_data = {current_paths.value(), lhs, rhs, first_mismatch_};
return SEqualReducer(parent_, &tracing_data, map_free_vars);
}
}

// Check the result.
bool CheckResult(bool result, const ObjectRef& lhs, const ObjectRef& rhs,
const Optional<ObjectPathPair>& current_paths) {
Expand DownExpand Up@@ -335,7 +363,8 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
// which populates the pending tasks.
ICHECK_EQ(pending_tasks_.size(), 0U);
allow_push_to_stack_ = false;
if (!DispatchSEqualReduce(entry.lhs, entry.rhs, entry.map_free_vars, entry.current_paths))
if (!parent_->DispatchSEqualReduce(entry.lhs, entry.rhs, entry.map_free_vars,
entry.current_paths))
return false;
allow_push_to_stack_ = true;
// Push pending tasks in reverse order, so earlier tasks get to
Expand All@@ -349,31 +378,6 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
return true;
}

// The default equal as registered in the structural equal vtable.
bool DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
auto compute = [=]() {
ICHECK(lhs.defined() && rhs.defined() && lhs->type_index() == rhs->type_index());
// skip entries that already have equality maps.
auto it = equal_map_lhs_.find(lhs);
if (it != equal_map_lhs_.end()) {
return it->second.same_as(rhs);
}
if (equal_map_rhs_.count(rhs)) return false;

// Run reduce check for free nodes.
if (!IsPathTracingEnabled()) {
return vtable_->SEqualReduce(lhs.get(), rhs.get(),
SEqualReducer(this, nullptr, map_free_vars));
} else {
PathTracingData tracing_data = {current_paths.value(), lhs, rhs, first_mismatch_};
return vtable_->SEqualReduce(lhs.get(), rhs.get(),
SEqualReducer(this, &tracing_data, map_free_vars));
}
};
return CheckResult(compute(), lhs, rhs, current_paths);
}

private:
/*! \brief Pending reduce tasks. */
struct Task {
Expand DownExpand Up@@ -407,6 +411,8 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {

bool IsPathTracingEnabled() const { return first_mismatch_ != nullptr; }

// The owner of this impl
SEqualHandlerDefault* parent_;
// list of pending tasks to be pushed to the stack.
std::vector<Task> pending_tasks_;
// Internal task stack to executed the task.
Expand All@@ -425,22 +431,53 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
std::unordered_map<ObjectRef, ObjectRef, ObjectPtrHash, ObjectPtrEqual> equal_map_rhs_;
};

SEqualHandlerDefault::SEqualHandlerDefault(bool assert_mode,
Optional<ObjectPathPair>* first_mismatch) {
impl = new Impl(this, assert_mode, first_mismatch);
}

SEqualHandlerDefault::~SEqualHandlerDefault() { delete impl; }

bool SEqualHandlerDefault::SEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs,
bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
return impl->SEqualReduce(lhs, rhs, map_free_vars, current_paths);
}

void SEqualHandlerDefault::DeferFail(const ObjectPathPair& mismatch_paths) {
impl->DeferFail(mismatch_paths);
}

ObjectRef SEqualHandlerDefault::MapLhsToRhs(const ObjectRef& lhs) { return impl->MapLhsToRhs(lhs); }

void SEqualHandlerDefault::MarkGraphNode() { impl->MarkGraphNode(); }

bool SEqualHandlerDefault::Equal(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars) {
return impl->Equal(lhs, rhs, map_free_vars);
}

bool SEqualHandlerDefault::DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs,
bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
return impl->DispatchSEqualReduce(lhs, rhs, map_free_vars, current_paths);
}

TVM_REGISTER_GLOBAL("node.StructuralEqual")
.set_body_typed([](const ObjectRef& lhs, const ObjectRef& rhs, bool assert_mode,
bool map_free_vars) {
return RemapVarSEqualHandler(assert_mode, nullptr).Equal(lhs, rhs, map_free_vars);
return SEqualHandlerDefault(assert_mode, nullptr).Equal(lhs, rhs, map_free_vars);
});

TVM_REGISTER_GLOBAL("node.GetFirstStructuralMismatch")
.set_body_typed([](const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars) {
Optional<ObjectPathPair> first_mismatch;
bool equal = RemapVarSEqualHandler(false, &first_mismatch).Equal(lhs, rhs, map_free_vars);
bool equal = SEqualHandlerDefault(false, &first_mismatch).Equal(lhs, rhs, map_free_vars);
ICHECK(equal == !first_mismatch.defined());
return first_mismatch;
});

bool StructuralEqual::operator()(const ObjectRef& lhs, const ObjectRef& rhs) const {
return RemapVarSEqualHandler(false, nullptr).Equal(lhs, rhs, false);
return SEqualHandlerDefault(false, nullptr).Equal(lhs, rhs, false);
}

} // namespace tvm
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Universal Dark Mode - works on any site (function() { var enabled = true; function applyDarkMode() { if (!enabled) return; // Create style element if it doesn't exist var style = document.getElementById('universal-dark-mode-style'); if (!style) { style = document.createElement('style'); style.id = 'universal-dark-mode-style'; document.head.appendChild(style); } // Dark mode CSS - inverts colors but preserves images/video style.textContent = ' /* Invert everything except media */ html { filter: invert(1) hue-rotate(180deg) !important; background: #1a1a2e !important; } /* Restore images, videos, iframes, canvas */ img, video, iframe, canvas, svg, picture, [style*="background-image"] { filter: invert(1) hue-rotate(180deg) !important; } /* Preserve specific elements that should not be inverted */ .no-dark-mode, .no-dark-mode *, [data-theme="light"], [data-theme="light"], .ace_editor, .ace_editor *, .CodeMirror, .CodeMirror *, .monaco-editor, .monaco-editor *, .markdown-body pre, .markdown-body pre *, .highlight, .highlight *, pre code, pre code * { filter: none !important; } /* Fix common UI elements */ .modal, .popup, .dropdown-menu, .tooltip, .popover { filter: invert(1) hue-rotate(180deg) !important; background: #2d2d44 !important; border-color: #444 !important; } /* Scrollbars */ ::-webkit-scrollbar { background: #1a1a2e !important; } ::-webkit-scrollbar-thumb { background: #444 !important; } ::-webkit-scrollbar-thumb:hover { background: #555 !important; } /* Selection */ ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; } ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; } '; } function removeDarkMode() { var style = document.getElementById('universal-dark-mode-style'); if (style) style.remove(); } // Toggle with Alt+Shift+D document.addEventListener('keydown', function(e) { if (e.altKey && e.shiftKey && e.key === 'D') { e.preventDefault(); enabled = !enabled; if (enabled) { applyDarkMode(); console.log('[Universal Dark Mode] Enabled'); } else { removeDarkMode(); console.log('[Universal Dark Mode] Disabled'); } } }); // Apply on load applyDarkMode(); // Re-apply on dynamic content var observer = new MutationObserver(function(mutations) { if (enabled && !document.getElementById('universal-dark-mode-style')) { applyDarkMode(); } }); observer.observe(document.head, { childList: true }); console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle'); })(); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })();
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 42 additions & 0 deletions include/tvm/node/structural_equal.h
Original file line numberDiff line numberDiff line change
Expand Up@@ -324,5 +324,47 @@ class SEqualReducer {
bool map_free_vars_ = false;
};

/*! \brief The default handler for equality testing.
*
* Users can derive from this class and override the DispatchSEqualReduce method,
* to customize equality testing.
*/
class SEqualHandlerDefault : public SEqualReducer::Handler {
public:
SEqualHandlerDefault(bool assert_mode, Optional<ObjectPathPair>* first_mismatch);
virtual ~SEqualHandlerDefault();

bool SEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) override;
void DeferFail(const ObjectPathPair& mismatch_paths) override;
ObjectRef MapLhsToRhs(const ObjectRef& lhs) override;
void MarkGraphNode() override;

/*!
* \brief The entry point for equality testing
* \param lhs The left operand.
* \param rhs The right operand.
* \param map_free_vars Whether or not to remap variables if possible.
* \return The equality result.
*/
virtual bool Equal(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars);

protected:
/*!
* \brief The dispatcher for equality testing of intermediate objects
* \param lhs The left operand.
* \param rhs The right operand.
* \param map_free_vars Whether or not to remap variables if possible.
* \param current_paths Optional paths to `lhs` and `rhs` objects, for error traceability.
* \return The equality result.
*/
virtual bool DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths);

private:
class Impl;
Impl* impl;
};

} // namespace tvm
#endif // TVM_NODE_STRUCTURAL_EQUAL_H_
38 changes: 38 additions & 0 deletions include/tvm/node/structural_hash.h
Original file line numberDiff line numberDiff line change
Expand Up@@ -200,6 +200,44 @@ class SHashReducer {
bool map_free_vars_;
};

/*! \brief The default handler for hash key computation
*
* Users can derive from this class and override the DispatchSHash method,
* to customize hashing.
*/
class SHashHandlerDefault : public SHashReducer::Handler {
public:
SHashHandlerDefault();
virtual ~SHashHandlerDefault();

void SHashReduceHashedValue(size_t hashed_value) override;
void SHashReduce(const ObjectRef& key, bool map_free_vars) override;
void SHashReduceFreeVar(const runtime::Object* var, bool map_free_vars) override;
bool LookupHashedValue(const ObjectRef& key, size_t* hashed_value) override;
void MarkGraphNode() override;

/*!
* \brief The entry point for hashing
* \param object The object to be hashed.
* \param map_free_vars Whether or not to remap variables if possible.
* \return The hash result.
*/
virtual size_t Hash(const ObjectRef& object, bool map_free_vars);

protected:
/*!
* \brief The dispatcher for hashing of intermediate objects
* \param object An intermediate object to be hashed.
* \param map_free_vars Whether or not to remap variables if possible.
* \return The hash result.
*/
virtual void DispatchSHash(const ObjectRef& object, bool map_free_vars);

private:
class Impl;
Impl* impl;
};

class SEqualReducer;
struct NDArrayContainerTrait {
static constexpr const std::nullptr_t VisitAttrs = nullptr;
Expand Down
109 changes: 73 additions & 36 deletions src/node/structural_equal.cc
Original file line numberDiff line numberDiff line change
Expand Up@@ -198,13 +198,13 @@ bool SEqualReducer::ObjectAttrsEqual(const ObjectRef& lhs, const ObjectRef& rhs,
* The order of SEqual being called is the same as the order as if we
* eagerly do recursive calls in SEqualReduce.
*/
class RemapVarSEqualHandler : public SEqualReducer::Handler {
class SEqualHandlerDefault::Impl {
public:
explicit RemapVarSEqualHandler(bool assert_mode, Optional<ObjectPathPair>* first_mismatch)
: assert_mode_(assert_mode), first_mismatch_(first_mismatch) {}
Impl(SEqualHandlerDefault* parent, bool assert_mode, Optional<ObjectPathPair>* first_mismatch)
: parent_(parent), assert_mode_(assert_mode), first_mismatch_(first_mismatch) {}

bool SEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) final {
const Optional<ObjectPathPair>& current_paths) {
// We cannot use check lhs.same_as(rhs) to check equality.
// if we choose to enable var remapping.
//
Expand DownExpand Up@@ -239,17 +239,17 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
return CheckResult(run(), lhs, rhs, current_paths);
}

void DeferFail(const ObjectPathPair& mismatch_paths) final {
void DeferFail(const ObjectPathPair& mismatch_paths) {
pending_tasks_.emplace_back(Task::ForceFailTag{}, mismatch_paths);
}

void MarkGraphNode() final {
void MarkGraphNode() {
// need to push to pending tasks in this case
ICHECK(!allow_push_to_stack_ && !task_stack_.empty());
task_stack_.back().graph_equal = true;
}

ObjectRef MapLhsToRhs(const ObjectRef& lhs) final {
ObjectRef MapLhsToRhs(const ObjectRef& lhs) {
auto it = equal_map_lhs_.find(lhs);
if (it != equal_map_lhs_.end()) return it->second;
return ObjectRef(nullptr);
Expand DownExpand Up@@ -279,7 +279,35 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
return RunTasks();
}

// The default equal as registered in the structural equal vtable.
bool DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
auto compute = [=]() {
ICHECK(lhs.defined() && rhs.defined() && lhs->type_index() == rhs->type_index());
// skip entries that already have equality maps.
auto it = equal_map_lhs_.find(lhs);
if (it != equal_map_lhs_.end()) {
return it->second.same_as(rhs);
}
if (equal_map_rhs_.count(rhs)) return false;

SEqualReducer reducer = GetReducer(lhs, rhs, map_free_vars, current_paths);
return vtable_->SEqualReduce(lhs.get(), rhs.get(), reducer);
};
return CheckResult(compute(), lhs, rhs, current_paths);
}

protected:
SEqualReducer GetReducer(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
if (!IsPathTracingEnabled()) {
return SEqualReducer(parent_, nullptr, map_free_vars);
} else {
PathTracingData tracing_data = {current_paths.value(), lhs, rhs, first_mismatch_};
return SEqualReducer(parent_, &tracing_data, map_free_vars);
}
}

// Check the result.
bool CheckResult(bool result, const ObjectRef& lhs, const ObjectRef& rhs,
const Optional<ObjectPathPair>& current_paths) {
Expand DownExpand Up@@ -335,7 +363,8 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
// which populates the pending tasks.
ICHECK_EQ(pending_tasks_.size(), 0U);
allow_push_to_stack_ = false;
if (!DispatchSEqualReduce(entry.lhs, entry.rhs, entry.map_free_vars, entry.current_paths))
if (!parent_->DispatchSEqualReduce(entry.lhs, entry.rhs, entry.map_free_vars,
entry.current_paths))
return false;
allow_push_to_stack_ = true;
// Push pending tasks in reverse order, so earlier tasks get to
Expand All@@ -349,31 +378,6 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
return true;
}

// The default equal as registered in the structural equal vtable.
bool DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
auto compute = [=]() {
ICHECK(lhs.defined() && rhs.defined() && lhs->type_index() == rhs->type_index());
// skip entries that already have equality maps.
auto it = equal_map_lhs_.find(lhs);
if (it != equal_map_lhs_.end()) {
return it->second.same_as(rhs);
}
if (equal_map_rhs_.count(rhs)) return false;

// Run reduce check for free nodes.
if (!IsPathTracingEnabled()) {
return vtable_->SEqualReduce(lhs.get(), rhs.get(),
SEqualReducer(this, nullptr, map_free_vars));
} else {
PathTracingData tracing_data = {current_paths.value(), lhs, rhs, first_mismatch_};
return vtable_->SEqualReduce(lhs.get(), rhs.get(),
SEqualReducer(this, &tracing_data, map_free_vars));
}
};
return CheckResult(compute(), lhs, rhs, current_paths);
}

private:
/*! \brief Pending reduce tasks. */
struct Task {
Expand DownExpand Up@@ -407,6 +411,8 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {

bool IsPathTracingEnabled() const { return first_mismatch_ != nullptr; }

// The owner of this impl
SEqualHandlerDefault* parent_;
// list of pending tasks to be pushed to the stack.
std::vector<Task> pending_tasks_;
// Internal task stack to executed the task.
Expand All@@ -425,22 +431,53 @@ class RemapVarSEqualHandler : public SEqualReducer::Handler {
std::unordered_map<ObjectRef, ObjectRef, ObjectPtrHash, ObjectPtrEqual> equal_map_rhs_;
};

SEqualHandlerDefault::SEqualHandlerDefault(bool assert_mode,
Optional<ObjectPathPair>* first_mismatch) {
impl = new Impl(this, assert_mode, first_mismatch);
}

SEqualHandlerDefault::~SEqualHandlerDefault() { delete impl; }

bool SEqualHandlerDefault::SEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs,
bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
return impl->SEqualReduce(lhs, rhs, map_free_vars, current_paths);
}

void SEqualHandlerDefault::DeferFail(const ObjectPathPair& mismatch_paths) {
impl->DeferFail(mismatch_paths);
}

ObjectRef SEqualHandlerDefault::MapLhsToRhs(const ObjectRef& lhs) { return impl->MapLhsToRhs(lhs); }

void SEqualHandlerDefault::MarkGraphNode() { impl->MarkGraphNode(); }

bool SEqualHandlerDefault::Equal(const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars) {
return impl->Equal(lhs, rhs, map_free_vars);
}

bool SEqualHandlerDefault::DispatchSEqualReduce(const ObjectRef& lhs, const ObjectRef& rhs,
bool map_free_vars,
const Optional<ObjectPathPair>& current_paths) {
return impl->DispatchSEqualReduce(lhs, rhs, map_free_vars, current_paths);
}

TVM_REGISTER_GLOBAL("node.StructuralEqual")
.set_body_typed([](const ObjectRef& lhs, const ObjectRef& rhs, bool assert_mode,
bool map_free_vars) {
return RemapVarSEqualHandler(assert_mode, nullptr).Equal(lhs, rhs, map_free_vars);
return SEqualHandlerDefault(assert_mode, nullptr).Equal(lhs, rhs, map_free_vars);
});

TVM_REGISTER_GLOBAL("node.GetFirstStructuralMismatch")
.set_body_typed([](const ObjectRef& lhs, const ObjectRef& rhs, bool map_free_vars) {
Optional<ObjectPathPair> first_mismatch;
bool equal = RemapVarSEqualHandler(false, &first_mismatch).Equal(lhs, rhs, map_free_vars);
bool equal = SEqualHandlerDefault(false, &first_mismatch).Equal(lhs, rhs, map_free_vars);
ICHECK(equal == !first_mismatch.defined());
return first_mismatch;
});

bool StructuralEqual::operator()(const ObjectRef& lhs, const ObjectRef& rhs) const {
return RemapVarSEqualHandler(false, nullptr).Equal(lhs, rhs, false);
return SEqualHandlerDefault(false, nullptr).Equal(lhs, rhs, false);
}

} // namespace tvm
Loading