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
6 changes: 4 additions & 2 deletions python/tvm/contrib/debugger/debug_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -272,8 +272,10 @@ def debug_get_output(self, node, out=None):
node_index = node
else:
raise RuntimeError("Require node index or name only.")

self._debug_get_output(node_index, out)
if out:
self._debug_get_output(node_index, out)
return out
return self._debug_get_output(node_index)

# pylint: disable=arguments-differ
def run(
Expand Down
23 changes: 21 additions & 2 deletions src/runtime/graph_executor/debug/graph_executor_debug.cc
Original file line number Diff line number Diff line change
Expand Up @@ -197,10 +197,17 @@ PackedFunc GraphExecutorDebug::GetFunction(const String& name,
// return member functions during query.
if (name == "debug_get_output") {
return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) {
int args0 = -1;
if (String::CanConvertFrom(args[0])) {
this->DebugGetNodeOutput(this->GetNodeIndex(args[0]), args[1]);
args0 = this->GetNodeIndex(args[0]);
} else {
this->DebugGetNodeOutput(args[0], args[1]);
args0 = args[0];
}

if (args.num_args == 2) {
this->DebugGetNodeOutput(args0, args[1]);
} else {
*rv = this->DebugGetNodeOutput(args0);
}
});
} else if (name == "execute_node") {
Expand Down Expand Up @@ -325,6 +332,18 @@ void GraphExecutorDebug::DebugGetNodeOutput(int index, DLTensor* data_out) {
data_entry_[eid].CopyTo(data_out);
}

NDArray GraphExecutorDebug::DebugGetNodeOutput(int index) {
ICHECK_LT(static_cast<size_t>(index), op_execs_.size());
uint32_t eid = index;

for (size_t i = 0; i < op_execs_.size(); ++i) {
if (op_execs_[i]) op_execs_[i]();
if (static_cast<int>(i) == index) break;
}

return data_entry_[eid];
}

NDArray GraphExecutorDebug::GetNodeOutput(int node, int out_ind) {
ICHECK_EQ(node, last_executed_node_);
ICHECK_LT(entry_id(node, out_ind), data_entry_.size());
Expand Down
12 changes: 12 additions & 0 deletions src/runtime/graph_executor/debug/graph_executor_debug.h
Original file line number Diff line number Diff line change
Expand Up @@ -122,6 +122,18 @@ class GraphExecutorDebug : public GraphExecutor {
*/
void DebugGetNodeOutput(int index, DLTensor* data_out);

/*!
* \brief return output of index-th node.
*
* This method will do a partial run of the graph
* from begining up to the index-th node and return output of index-th node.
* This is costly operation and suggest to use only for debug porpose.
*
* \param index: The index of the node.
*
*/
NDArray DebugGetNodeOutput(int index);

/*!
* \brief Profile execution time of the module.
*
Expand Down