From 267644076bb8f4cc8744825c8ca079685579c80d Mon Sep 17 00:00:00 2001 From: Junru Shao Date: Thu, 14 Dec 2023 16:24:31 +0000 Subject: [PATCH 1/5] fix --- src/runtime/ndarray.cc | 21 +++++++++++++++++++++ src/target/source/codegen_c.cc | 34 +++++++++++++++++++++++++++------- 2 files changed, 48 insertions(+), 7 deletions(-) diff --git a/src/runtime/ndarray.cc b/src/runtime/ndarray.cc index b7153ab50f1f..fa6bf563be89 100644 --- a/src/runtime/ndarray.cc +++ b/src/runtime/ndarray.cc @@ -96,6 +96,27 @@ void ArrayCopyToBytes(const DLTensor* handle, void* data, size_t nbytes) { DeviceAPI::Get(handle->device)->StreamSync(handle->device, nullptr); } +TVM_REGISTER_GLOBAL("mlc.show_dltensor").set_body_typed([](NDArray arr) { + auto handle = arr.operator->(); + LOG(INFO) << "data: " << handle->data; + LOG(INFO) << "device:" << handle->device.device_type << " " << handle->device.device_id; + LOG(INFO) << "ndim: " << handle->ndim; + LOG(INFO) << DLDataType2String(handle->dtype); + LOG(INFO) << "shape:"; + for (int i = 0; i < handle->ndim; ++i) { + LOG(INFO) << handle->shape[i]; + } + LOG(INFO) << "strides:"; + if (handle->strides == nullptr) { + LOG(INFO) << "nullptr"; + } else { + for (int i = 0; i < handle->ndim; ++i) { + LOG(INFO) << handle->strides[i]; + } + } + LOG(INFO) << "byte_offset: " << handle->byte_offset; +}); + struct NDArray::Internal { // Default deleter for the container static void DefaultDeleter(Object* ptr_obj) { diff --git a/src/target/source/codegen_c.cc b/src/target/source/codegen_c.cc index 0ff0531b5c20..7cfe6ae06c79 100644 --- a/src/target/source/codegen_c.cc +++ b/src/target/source/codegen_c.cc @@ -631,13 +631,33 @@ void CodeGenC::VisitExpr_(const CallNode* op, std::ostream& os) { // NOLINT(*) } else if (op->op.same_as(builtin::shift_right())) { PrintBinaryIntrinsic(op, " >> ", os, this); } else if (op->op.same_as(builtin::if_then_else())) { - os << "("; - PrintExpr(op->args[0], os); - os << " ? "; - PrintExpr(op->args[1], os); - os << " : "; - PrintExpr(op->args[2], os); - os << ")"; + // conditional that skips eval if cond evals to false + std::string result = name_supply_->FreshName("condval"); + std::string cond = PrintExpr(op->args[0]); + this->PrintIndent(); + PrintType(op->dtype, this->stream); + this->stream << " " << result << ";\n"; + this->PrintIndent(); + this->stream << "if (" << cond << ") {\n"; + { + int then_scope = this->BeginScope(); + std::string true_val = PrintExpr(op->args[1]); + this->PrintIndent(); + this->stream << result << " = " << true_val << ";\n"; + this->EndScope(then_scope); + this->PrintIndent(); + this->stream << "} else {\n"; + } + { + int else_scope = this->BeginScope(); + std::string false_val = PrintExpr(op->args[2]); + this->PrintIndent(); + this->stream << result << " = " << false_val << ";\n"; + this->EndScope(else_scope); + this->PrintIndent(); + this->stream << "}\n"; + } + os << result; } else if (op->op.same_as(builtin::address_of())) { const BufferLoadNode* load = op->args[0].as(); ICHECK(op->args.size() == 1 && load); From 74eccefe1fb15496923b6f2a87069f46dfaaca85 Mon Sep 17 00:00:00 2001 From: Junru Shao Date: Thu, 14 Dec 2023 17:38:28 +0000 Subject: [PATCH 2/5] lint --- src/runtime/ndarray.cc | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/runtime/ndarray.cc b/src/runtime/ndarray.cc index fa6bf563be89..83200e089117 100644 --- a/src/runtime/ndarray.cc +++ b/src/runtime/ndarray.cc @@ -99,7 +99,7 @@ void ArrayCopyToBytes(const DLTensor* handle, void* data, size_t nbytes) { TVM_REGISTER_GLOBAL("mlc.show_dltensor").set_body_typed([](NDArray arr) { auto handle = arr.operator->(); LOG(INFO) << "data: " << handle->data; - LOG(INFO) << "device:" << handle->device.device_type << " " << handle->device.device_id; + LOG(INFO) << "device:" << handle->device.device_type << " " << handle->device.device_id; LOG(INFO) << "ndim: " << handle->ndim; LOG(INFO) << DLDataType2String(handle->dtype); LOG(INFO) << "shape:"; From a7c3ef1db9b1befa0dd9a08474fe12113ccfd03a Mon Sep 17 00:00:00 2001 From: spectrometerHBH Date: Tue, 2 Jan 2024 14:30:41 -0500 Subject: [PATCH 3/5] fix while loop --- src/target/source/codegen_c.cc | 5 ++++- src/tir/transforms/thread_storage_sync.cc | 4 ++-- 2 files changed, 6 insertions(+), 3 deletions(-) diff --git a/src/target/source/codegen_c.cc b/src/target/source/codegen_c.cc index 7cfe6ae06c79..6c22b69d1908 100644 --- a/src/target/source/codegen_c.cc +++ b/src/target/source/codegen_c.cc @@ -1038,8 +1038,11 @@ void CodeGenC::VisitStmt_(const ForNode* op) { void CodeGenC::VisitStmt_(const WhileNode* op) { PrintIndent(); - stream << "while (" << PrintExpr(op->condition) << ") {\n"; + stream << "while (1) {\n"; int while_scope = BeginScope(); + std::string cond = PrintExpr(op->condition); + PrintIndent(); + stream << "if (!(" << cond << ")) { break; }\n"; PrintStmt(op->body); this->EndScope(while_scope); PrintIndent(); diff --git a/src/tir/transforms/thread_storage_sync.cc b/src/tir/transforms/thread_storage_sync.cc index d92986e51a9c..b2de24c9674e 100644 --- a/src/tir/transforms/thread_storage_sync.cc +++ b/src/tir/transforms/thread_storage_sync.cc @@ -113,7 +113,7 @@ class ThreadSyncPlanner : public StorageAccessVisitor { } } if (sync_before_stmt) { - ICHECK_EQ(condition_counter(), 0) << "Cannot insert syncs inside condition"; + // ICHECK_EQ(condition_counter(), 0) << "Cannot insert syncs inside condition"; syncs_inserted_.insert(s.stmt); } } @@ -140,7 +140,7 @@ class ThreadSyncPlanner : public StorageAccessVisitor { } } if (sync_before_stmt) { - ICHECK_EQ(condition_counter(), 0) << "Cannot insert syncs inside condition"; + // ICHECK_EQ(condition_counter(), 0) << "Cannot insert syncs inside condition"; syncs_inserted_.insert(s.stmt); break; } From b207b493099011be44c3ac0a60b919812dde495d Mon Sep 17 00:00:00 2001 From: spectrometerHBH Date: Tue, 2 Jan 2024 14:35:12 -0500 Subject: [PATCH 4/5] clean --- src/runtime/ndarray.cc | 21 --------------------- 1 file changed, 21 deletions(-) diff --git a/src/runtime/ndarray.cc b/src/runtime/ndarray.cc index 83200e089117..b7153ab50f1f 100644 --- a/src/runtime/ndarray.cc +++ b/src/runtime/ndarray.cc @@ -96,27 +96,6 @@ void ArrayCopyToBytes(const DLTensor* handle, void* data, size_t nbytes) { DeviceAPI::Get(handle->device)->StreamSync(handle->device, nullptr); } -TVM_REGISTER_GLOBAL("mlc.show_dltensor").set_body_typed([](NDArray arr) { - auto handle = arr.operator->(); - LOG(INFO) << "data: " << handle->data; - LOG(INFO) << "device:" << handle->device.device_type << " " << handle->device.device_id; - LOG(INFO) << "ndim: " << handle->ndim; - LOG(INFO) << DLDataType2String(handle->dtype); - LOG(INFO) << "shape:"; - for (int i = 0; i < handle->ndim; ++i) { - LOG(INFO) << handle->shape[i]; - } - LOG(INFO) << "strides:"; - if (handle->strides == nullptr) { - LOG(INFO) << "nullptr"; - } else { - for (int i = 0; i < handle->ndim; ++i) { - LOG(INFO) << handle->strides[i]; - } - } - LOG(INFO) << "byte_offset: " << handle->byte_offset; -}); - struct NDArray::Internal { // Default deleter for the container static void DefaultDeleter(Object* ptr_obj) { From d2a0f548a0417b5fdf3ed0e3927dc7a697d5b47b Mon Sep 17 00:00:00 2001 From: spectrometerHBH Date: Wed, 3 Jan 2024 14:02:08 -0500 Subject: [PATCH 5/5] clean --- src/tir/transforms/thread_storage_sync.cc | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/tir/transforms/thread_storage_sync.cc b/src/tir/transforms/thread_storage_sync.cc index b2de24c9674e..d92986e51a9c 100644 --- a/src/tir/transforms/thread_storage_sync.cc +++ b/src/tir/transforms/thread_storage_sync.cc @@ -113,7 +113,7 @@ class ThreadSyncPlanner : public StorageAccessVisitor { } } if (sync_before_stmt) { - // ICHECK_EQ(condition_counter(), 0) << "Cannot insert syncs inside condition"; + ICHECK_EQ(condition_counter(), 0) << "Cannot insert syncs inside condition"; syncs_inserted_.insert(s.stmt); } } @@ -140,7 +140,7 @@ class ThreadSyncPlanner : public StorageAccessVisitor { } } if (sync_before_stmt) { - // ICHECK_EQ(condition_counter(), 0) << "Cannot insert syncs inside condition"; + ICHECK_EQ(condition_counter(), 0) << "Cannot insert syncs inside condition"; syncs_inserted_.insert(s.stmt); break; }