Skip to content

Commit 31f8c8a

Browse files
authored
[Runtime][Vulkan] Add timer support for the Vulkan backend (#20471)
## Motivation The Vulkan backend has no device-specific `Timer` implementation, so profiling falls back to the default CPU wall-clock timer. That timer measures host-side submission and synchronization, not actual GPU execution, so profiling results on Vulkan devices are inaccurate. ## Changes - Add `VulkanTimerNode` (`vulkan_timer.h` / `vulkan_timer.cc`), which measures GPU execution time using Vulkan timestamp queries (`vkCmdWriteTimestamp`) recorded on the device's compute stream. - Store `timestampPeriod` and the compute queue family's `timestampValidBits` in `VulkanDeviceProperties`. They are used to convert ticks to nanoseconds and to mask invalid upper bits of the timestamp values. - Check at timer construction that the compute queue supports timestamp queries (`timestampValidBits > 0`) and fail with a clear error if it does not. - Register the Vulkan timer factory so `Timer::Start(dev)` picks it up automatically for Vulkan devices.
1 parent 135dea7 commit 31f8c8a

5 files changed

Lines changed: 215 additions & 0 deletions

File tree

‎src/backend/vulkan/runtime/vulkan_device.cc‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -165,6 +165,7 @@ VulkanDeviceProperties::VulkanDeviceProperties(const VulkanInstance& instance,
165165
max_shared_memory_per_block = properties.properties.limits.maxComputeSharedMemorySize;
166166
device_name = properties.properties.deviceName;
167167
driver_version = properties.properties.driverVersion;
168+
timestamp_period = properties.properties.limits.timestampPeriod;
168169

169170
if (device.HasExtension("VK_KHR_driver_properties")) {
170171
driver_name = driver.driverName;

‎src/backend/vulkan/runtime/vulkan_device.h‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -106,6 +106,7 @@ struct VulkanDeviceProperties {
106106
uint32_t driver_version{0};
107107
uint32_t vulkan_api_version{VK_API_VERSION_1_0};
108108
uint32_t max_spirv_version{0x10000};
109+
double timestamp_period{0};
109110
};
110111

111112
/*! \brief Handle to the Vulkan API's VkDevice

‎src/backend/vulkan/runtime/vulkan_device_api.cc‎

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,7 @@
2828
#include <utility>
2929

3030
#include "vulkan_common.h"
31+
#include "vulkan_timer.h"
3132

3233
namespace tvm {
3334
namespace runtime {
@@ -452,6 +453,18 @@ VulkanDevice& VulkanDeviceAPI::device(size_t device_id) {
452453
return const_cast<VulkanDevice&>(const_cast<const VulkanDeviceAPI*>(this)->device(device_id));
453454
}
454455

456+
TVM_FFI_STATIC_INIT_BLOCK() {
457+
namespace refl = tvm::ffi::reflection;
458+
refl::GlobalDef().def("profiling.timer.vulkan",
459+
[](Device dev) { return Timer(ffi::make_object<VulkanTimerNode>(dev)); });
460+
}
461+
462+
TVM_FFI_STATIC_INIT_BLOCK() {
463+
namespace refl = tvm::ffi::reflection;
464+
refl::GlobalDef().def("runtime.timer.vulkan",
465+
[](Device dev) { return Timer(ffi::make_object<VulkanTimerNode>(dev)); });
466+
}
467+
455468
TVM_FFI_STATIC_INIT_BLOCK() {
456469
namespace refl = tvm::ffi::reflection;
457470
refl::GlobalDef()
Lines changed: 100 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,100 @@
1+
/*
2+
* Licensed to the Apache Software Foundation (ASF) under one
3+
* or more contributor license agreements. See the NOTICE file
4+
* distributed with this work for additional information
5+
* regarding copyright ownership. The ASF licenses this file
6+
* to you under the Apache License, Version 2.0 (the
7+
* "License"); you may not use this file except in compliance
8+
* with the License. You may obtain a copy of the License at
9+
*
10+
* http://www.apache.org/licenses/LICENSE-2.0
11+
*
12+
* Unless required by applicable law or agreed to in writing,
13+
* software distributed under the License is distributed on an
14+
* "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
15+
* KIND, either express or implied. See the License for the
16+
* specific language governing permissions and limitations
17+
* under the License.
18+
*/
19+
20+
#include "vulkan_timer.h"
21+
22+
#include <tvm/runtime/logging.h>
23+
24+
#include "vulkan_device_api.h"
25+
26+
namespace tvm {
27+
namespace runtime {
28+
namespace vulkan {
29+
30+
VulkanTimerNode::VulkanTimerNode(Device dev) : dev_(dev) {
31+
// Get the Vulkan device and stream
32+
auto& vk_dev = VulkanDeviceAPI::Global()->device(dev_.device_id);
33+
stream_ = &vk_dev.ThreadLocalStream();
34+
device_ = vk_dev;
35+
36+
// Retrieve the timestamp period from device properties
37+
timestamp_period_ = vk_dev.device_properties.timestamp_period;
38+
TVM_FFI_ICHECK_GT(timestamp_period_, 0) << "Vulkan device does not support timestamp queries.";
39+
40+
CreateQueryPool();
41+
}
42+
43+
VulkanTimerNode::~VulkanTimerNode() { Cleanup(); }
44+
45+
void VulkanTimerNode::CreateQueryPool() {
46+
VkQueryPoolCreateInfo query_pool_info{};
47+
query_pool_info.sType = VK_STRUCTURE_TYPE_QUERY_POOL_CREATE_INFO;
48+
query_pool_info.queryType = VK_QUERY_TYPE_TIMESTAMP;
49+
query_pool_info.queryCount = 2;
50+
51+
VkResult res = vkCreateQueryPool(device_, &query_pool_info, nullptr, &query_pool_);
52+
TVM_FFI_ICHECK(res == VK_SUCCESS) << "Failed to create Vulkan query pool.";
53+
}
54+
55+
void VulkanTimerNode::Start() {
56+
stream_->Launch([this](VulkanStreamState* state) {
57+
// Reset the query pool before writing timestamps
58+
vkCmdResetQueryPool(state->cmd_buffer_, query_pool_, start_query_, 2);
59+
vkCmdWriteTimestamp(state->cmd_buffer_, VK_PIPELINE_STAGE_TOP_OF_PIPE_BIT, query_pool_,
60+
start_query_);
61+
});
62+
}
63+
64+
void VulkanTimerNode::Stop() {
65+
stream_->Launch([this](VulkanStreamState* state) {
66+
vkCmdWriteTimestamp(state->cmd_buffer_, VK_PIPELINE_STAGE_BOTTOM_OF_PIPE_BIT, query_pool_,
67+
end_query_);
68+
});
69+
70+
// Ensure GPU has finished writing timestamps before collecting them
71+
stream_->Synchronize();
72+
CollectTimestamps();
73+
}
74+
75+
int64_t VulkanTimerNode::SyncAndGetElapsedNanos() { return duration_; }
76+
77+
void VulkanTimerNode::CollectTimestamps() {
78+
uint64_t timestamps[2] = {0};
79+
80+
VkResult result =
81+
vkGetQueryPoolResults(device_, query_pool_, 0, 2, sizeof(timestamps), timestamps,
82+
sizeof(uint64_t), VK_QUERY_RESULT_64_BIT | VK_QUERY_RESULT_WAIT_BIT);
83+
84+
TVM_FFI_ICHECK(result == VK_SUCCESS) << "Failed to get Vulkan query pool results.";
85+
86+
// Calculate the duration in nanoseconds
87+
uint64_t diff = timestamps[1] - timestamps[0];
88+
duration_ = static_cast<int64_t>(static_cast<double>(diff) * timestamp_period_);
89+
}
90+
91+
void VulkanTimerNode::Cleanup() {
92+
if (query_pool_ != VK_NULL_HANDLE) {
93+
vkDestroyQueryPool(device_, query_pool_, nullptr);
94+
query_pool_ = VK_NULL_HANDLE;
95+
}
96+
}
97+
98+
} // namespace vulkan
99+
} // namespace runtime
100+
} // namespace tvm
Lines changed: 100 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,100 @@
1+
/*
2+
* Licensed to the Apache Software Foundation (ASF) under one
3+
* or more contributor license agreements. See the NOTICE file
4+
* distributed with this work for additional information
5+
* regarding copyright ownership. The ASF licenses this file
6+
* to you under the Apache License, Version 2.0 (the
7+
* "License"); you may not use this file except in compliance
8+
* with the License. You may obtain a copy of the License at
9+
*
10+
* http://www.apache.org/licenses/LICENSE-2.0
11+
*
12+
* Unless required by applicable law or agreed to in writing,
13+
* software distributed under the License is distributed on an
14+
* "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
15+
* KIND, either express or implied. See the License for the
16+
* specific language governing permissions and limitations
17+
* under the License.
18+
*/
19+
20+
#ifndef TVM_RUNTIME_VULKAN_VULKAN_TIMER_H_
21+
#define TVM_RUNTIME_VULKAN_VULKAN_TIMER_H_
22+
23+
#include <tvm/runtime/timer.h>
24+
#include <vulkan/vulkan.h>
25+
26+
#include "vulkan_device.h"
27+
#include "vulkan_stream.h"
28+
29+
namespace tvm {
30+
namespace runtime {
31+
namespace vulkan {
32+
33+
/*!
34+
* \brief Timer node for measuring GPU execution time using Vulkan.
35+
*
36+
* This class uses Vulkan timestamp queries to measure the time taken
37+
* by GPU operations between `Start()` and `Stop()` calls.
38+
*/
39+
class VulkanTimerNode : public TimerNode {
40+
public:
41+
/*!
42+
* \brief Constructs a VulkanTimerNode for the specified device.
43+
* \param dev The TVM device to be used for timing.
44+
*/
45+
explicit VulkanTimerNode(Device dev);
46+
47+
/*!
48+
* \brief Destructor to clean up Vulkan resources.
49+
*/
50+
~VulkanTimerNode() override;
51+
52+
/*!
53+
* \brief Starts the timer by recording a timestamp.
54+
*/
55+
void Start() override;
56+
57+
/*!
58+
* \brief Stops the timer by recording another timestamp.
59+
*/
60+
void Stop() override;
61+
62+
/*!
63+
* \brief Retrieves the elapsed time in nanoseconds.
64+
* \return The elapsed time in nanoseconds between Start and Stop.
65+
*/
66+
int64_t SyncAndGetElapsedNanos() override;
67+
68+
TVM_FFI_DECLARE_OBJECT_INFO_FINAL("runtime.vulkan.VulkanTimerNode", VulkanTimerNode, TimerNode);
69+
70+
private:
71+
Device dev_; ///< The TVM device being used.
72+
VkDevice device_{VK_NULL_HANDLE}; ///< The Vulkan device handle.
73+
VulkanStream* stream_{nullptr}; ///< The Vulkan stream for command buffer management.
74+
VkQueryPool query_pool_{VK_NULL_HANDLE}; ///< The Vulkan query pool for timestamp queries.
75+
double timestamp_period_; ///< The period (in nanoseconds) for each timestamp tick.
76+
uint32_t start_query_ = 0; ///< The index for the start timestamp query.
77+
uint32_t end_query_ = 1; ///< The index for the end timestamp query.
78+
int64_t duration_ = 0; ///< The measured duration in nanoseconds.
79+
80+
/*!
81+
* \brief Creates a Vulkan query pool for timestamp queries.
82+
*/
83+
void CreateQueryPool();
84+
85+
/*!
86+
* \brief Collects timestamps and calculates the duration.
87+
*/
88+
void CollectTimestamps();
89+
90+
/*!
91+
* \brief Cleans up the Vulkan query pool.
92+
*/
93+
void Cleanup();
94+
};
95+
96+
} // namespace vulkan
97+
} // namespace runtime
98+
} // namespace tvm
99+
100+
#endif // TVM_RUNTIME_VULKAN_VULKAN_TIMER_H_

0 commit comments

Comments
 (0)