From 14f238247bcd08b127b944b9f853288aa3b6f5ea Mon Sep 17 00:00:00 2001 From: baldurk Date: Wed, 8 Jan 2025 13:46:58 +0000 Subject: [PATCH] Add unit tests to GPUVA tracker and keep implementation private --- renderdoc/core/gpu_address_range_tracker.cpp | 648 ++++++++++++++++--- renderdoc/core/gpu_address_range_tracker.h | 48 +- renderdoc/driver/d3d12/d3d12_debug.cpp | 4 +- renderdoc/driver/d3d12/d3d12_debug.h | 2 +- renderdoc/driver/d3d12/d3d12_device.cpp | 3 +- renderdoc/driver/d3d12/d3d12_manager.cpp | 11 +- renderdoc/driver/d3d12/d3d12_manager.h | 2 +- renderdoc/driver/d3d12/d3d12_resources.cpp | 19 +- 8 files changed, 611 insertions(+), 126 deletions(-) diff --git a/renderdoc/core/gpu_address_range_tracker.cpp b/renderdoc/core/gpu_address_range_tracker.cpp index 49f77e855..5c2a9cb21 100644 --- a/renderdoc/core/gpu_address_range_tracker.cpp +++ b/renderdoc/core/gpu_address_range_tracker.cpp @@ -23,42 +23,156 @@ ******************************************************************************/ #include "core/gpu_address_range_tracker.h" +#include "api/replay/replay_enums.h" +#include "common/formatting.h" +#include "core/settings.h" void GPUAddressRangeTracker::AddTo(const GPUAddressRange &range) { SCOPED_WRITELOCK(addressLock); - auto it = std::lower_bound(addresses.begin(), addresses.end(), range.start); - // for resources with identical start addresses, sort them by realEnd - while(it != addresses.end() && it->start == range.start && it->realEnd < range.realEnd) - ++it; + // insert ranges ordered by start first, then by size. Ranges with different sizes starting at the + // same point will be ordered such that the last one is largest - addresses.insert(it - addresses.begin(), range); + // search for the range. This will return the largest range which starts before or at this address + size_t idx = FindLastRangeBeforeOrAtAddress(range.start); + + // if we search for an address that's past the end of the last range, we'll return that index. The + // only case where we return no valid index is if the address is before the first range - so + // insert ours at the start of the list and return + if(idx == ~0U) + { + addresses.insert(0, range); + return; + } + + // if the range found doesn't start at the same point as us, insert immediately so we preserve the + // sorting by range start + if(addresses[idx].start != range.start) + { + addresses.insert(idx + 1, range); + return; + } + + // we get here if the range starts at the same point as us, so we need to sort by size. + // if we are smaller than the found range, move backwards to insert before it. Keep going as long + // as we're looking at ranges that start at the same address and are larger than us + while(addresses[idx].start == range.start && addresses[idx].realEnd > range.realEnd) + { + // we could be smaller than the very first range in the list. If that's the case, insert at 0 and return now + if(idx == 0) + { + addresses.insert(0, range); + return; + } + + // otherwise move backwards, to insert before the current range + idx--; + } + + // insert after the idx we arrived at, which is the first range either starting before, or that is smaller than us + addresses.insert(idx + 1, range); } void GPUAddressRangeTracker::RemoveFrom(const GPUAddressRange &range) { { SCOPED_WRITELOCK(addressLock); - size_t i = std::lower_bound(addresses.begin(), addresses.end(), range.start) - addresses.begin(); - // there might be multiple buffers with the same range start, find the exact range for this - // buffer - while(i < addresses.size() && addresses[i].start == range.start) + // search for the range. This will return the largest range which starts before or at this address + size_t idx = FindLastRangeBeforeOrAtAddress(range.start); + + if(idx != ~0U) { - if(addresses[i].id == range.id) + // there might be multiple buffers with the same range start, find the exact range for this + // buffer. We only have to search backwards because we returned the largest (aka last) range before this address + while(addresses[idx].start == range.start) { - addresses.erase(i); - return; - } + if(addresses[idx].id == range.id) + { + addresses.erase(idx); + return; + } - ++i; + if(idx == 0) + break; + + --idx; + } } } - RDCERR("Couldn't find matching range to remove for %s", ToStr(range.id).c_str()); + // used only so the tests can EXPECT_ERROR() + RDResult err; + SET_ERROR_RESULT(err, ResultCode::InternalError, "Couldn't find matching range to remove for %s", + ToStr(range.id).c_str()); + (void)err; } +void GPUAddressRangeTracker::Clear() +{ + SCOPED_WRITELOCK(addressLock); + addresses.clear(); +} + +rdcarray GPUAddressRangeTracker::GetAddresses() +{ + SCOPED_READLOCK(addressLock); + return addresses; +} + +rdcarray GPUAddressRangeTracker::GetIDs() +{ + rdcarray ret; + ret.reserve(addresses.size()); + + { + SCOPED_READLOCK(addressLock); + for(size_t i = 0; i < addresses.size(); i++) + ret.push_back(addresses[i].id); + } + + return ret; +} + +size_t GPUAddressRangeTracker::FindLastRangeBeforeOrAtAddress(GPUAddressRange::Address addr) +{ + // the caller must lock. + + if(addresses.empty()) + return ~0U; + + // start looking at the whole range + size_t first = 0; + size_t count = addresses.size(); + + while(count > 1) + { + // look at the midpoint + size_t halfrange = count / 2; + size_t mid = first + halfrange; + + // if the midpoint is after our address, bisect down to the lower half and exclude the midpoint + if(addr < addresses[mid].start) + { + count = halfrange; + } + else + { + // midpoint is before or at our address, use upper half + first = mid; + count -= halfrange; + } + } + + // if first is 0 and the address range doesn't match, indicate that by returning ~0U + if(first == 0 && addr < addresses[first].start) + return ~0U; + + return first; +} + +template void GPUAddressRangeTracker::GetResIDFromAddr(GPUAddressRange::Address addr, ResourceId &id, uint64_t &offs) { @@ -73,64 +187,28 @@ void GPUAddressRangeTracker::GetResIDFromAddr(GPUAddressRange::Address addr, Res { SCOPED_READLOCK(addressLock); - auto it = std::lower_bound(addresses.begin(), addresses.end(), addr); - if(it == addresses.end()) + // search for the address. This will return the largest range which starts before or at this address + size_t idx = FindLastRangeBeforeOrAtAddress(addr); + + // ~0U is returned if the address is before the first range in our list. That means no match + if(idx == ~0U) return; - range = *it; - - // find the largest resource containing this address - not perfect but helps with trivially bad - // aliases where a tiny resource and a large resource are co-situated and the larger resource - // needs to be used for validity - while((it + 1) != addresses.end() && (it + 1)->start <= addr && (it + 1)->realEnd >= range.realEnd) - { - it++; - range = *it; - } - } - - if(addr < range.start || addr >= range.realEnd) - return; - - id = range.id; - offs = addr - range.start; -} - -void GPUAddressRangeTracker::GetResIDFromAddrAllowOutOfBounds(GPUAddressRange::Address addr, - ResourceId &id, uint64_t &offs) -{ - id = ResourceId(); - offs = 0; - - if(addr == 0) - return; - - GPUAddressRange range; - - { - SCOPED_READLOCK(addressLock); - - auto it = std::lower_bound(addresses.begin(), addresses.end(), addr); - if(it == addresses.end()) - return; - - range = *it; - - // find the largest resource containing this address - not perfect but helps with trivially bad - // aliases where a tiny resource and a large resource are co-situated and the larger resource - // needs to be used for validity - while((it + 1) != addresses.end() && (it + 1)->start <= addr && (it + 1)->realEnd >= range.realEnd) - { - it++; - range = *it; - } + // this range is already the largest before or at the address by virtue of our sorting and search + range = addresses[idx]; } if(addr < range.start) return; - // still enforce the OOB end on ranges - which is the remaining range in the backing store. - // Otherwise we could end up passing through invalid addresses stored in stale descriptors + // if OOB isn't allowed, check against real end + if(!allowOOB) + { + if(addr >= range.realEnd) + return; + } + + // always check against OOB end if(addr >= range.oobEnd) return; @@ -138,6 +216,11 @@ void GPUAddressRangeTracker::GetResIDFromAddrAllowOutOfBounds(GPUAddressRange::A offs = addr - range.start; } +template void GPUAddressRangeTracker::GetResIDFromAddr(GPUAddressRange::Address addr, + ResourceId &id, uint64_t &offs); +template void GPUAddressRangeTracker::GetResIDFromAddr(GPUAddressRange::Address addr, + ResourceId &id, uint64_t &offs); + void GPUAddressRangeTracker::GetResIDBoundForAddr(GPUAddressRange::Address addr, ResourceId &lower, GPUAddressRange::Address &lowerVA, ResourceId &upper, @@ -149,47 +232,432 @@ void GPUAddressRangeTracker::GetResIDBoundForAddr(GPUAddressRange::Address addr, if(addr == 0) return; - if(addresses.empty()) - return; - { SCOPED_READLOCK(addressLock); - auto it = std::lower_bound(addresses.begin(), addresses.end(), addr); - if(it == addresses.end()) - { - --it; + if(addresses.empty()) + return; - lower = it->id; - lowerVA = it->start; + size_t idx = FindLastRangeBeforeOrAtAddress(addr); + + // if the addr is before first known range, it's bounded on upper only + if(idx == ~0U) + { + upper = addresses[idx].id; + upperVA = addresses[idx].start; return; } - // find the last resource containing this address if there are multiple overlapping - while((it + 1) != addresses.end() && (it + 1)->start <= addr && (it + 1)->realEnd > addr) - { - it++; - } - - lower = it->id; - lowerVA = it->start; + lower = addresses[idx].id; + lowerVA = addresses[idx].start; // if this range contains the address exactly, return it as a tight bound - if(it->realEnd > addr) + if(addresses[idx].realEnd > addr) { - upper = it->id; - upperVA = it->realEnd; + upper = addresses[idx].id; + upperVA = addresses[idx].realEnd; + return; } // otherwise the address is past its end but before the next. Move one allocation along - we // already know that we picked the largest allocation that covers this address - ++it; + idx++; // if this wasn't the end, return the upper bound - if(it != addresses.end()) + if(idx < addresses.size()) { - upper = it->id; - upperVA = it->realEnd; + upper = addresses[idx].id; + upperVA = addresses[idx].realEnd; } } } + +#if ENABLED(ENABLE_UNIT_TESTS) + +#undef None +#undef Always + +#include "catch/catch.hpp" + +namespace TestIDs +{ +ResourceId a = ResourceIDGen::GetNewUniqueID(); +ResourceId b = ResourceIDGen::GetNewUniqueID(); +ResourceId c = ResourceIDGen::GetNewUniqueID(); +ResourceId d = ResourceIDGen::GetNewUniqueID(); +ResourceId e = ResourceIDGen::GetNewUniqueID(); +ResourceId f = ResourceIDGen::GetNewUniqueID(); +ResourceId g = ResourceIDGen::GetNewUniqueID(); +}; + +template <> +rdcstr DoStringise(const rdcpair &el) +{ + using namespace TestIDs; + + rdcarray ids = {a, b, c, d, e, f, g}; + rdcstr idname = "a"; + int idx = ids.indexOf(el.first); + if(idx >= 0) + idname[0] += (char)idx; + else if(el.first == ResourceId()) + idname[0] = '-'; + else + idname = "?"; + + return "{ " + idname + ", " + StringFormat::Fmt("%#x", el.second) + " }"; +} + +static GPUAddressRange MakeRange(ResourceId id, GPUAddressRange::Address addr, uint64_t size, + uint64_t oobPadding = 0) +{ + return { + addr, + addr + size, + addr + size + oobPadding, + id, + }; +} + +rdcpair make_idoffs(ResourceId a, uint64_t b) +{ + return {a, b}; +} + +TEST_CASE("Check GPUAddressRangeTracker", "[gpuaddr]") +{ + GPUAddressRangeTracker tracker; + + rdcpair none = make_idoffs(ResourceId(), 0ULL); + + using namespace TestIDs; + + SECTION("Basics") + { + tracker.AddTo(MakeRange(a, 0x1230000, 128)); + tracker.AddTo(MakeRange(b, 0x1250000, 128)); + + CHECK(tracker.GetResIDFromAddr(0) == none); + + CHECK(tracker.GetResIDFromAddr(0x1230000 - 1) == none); + CHECK(tracker.GetResIDFromAddr(0x1230000) == make_idoffs(a, 0ULL)); + CHECK(tracker.GetResIDFromAddr(0x1230001) == make_idoffs(a, 1ULL)); + CHECK(tracker.GetResIDFromAddr(0x1230000 + 127) == make_idoffs(a, 127ULL)); + CHECK(tracker.GetResIDFromAddr(0x1230000 + 128) == none); + + CHECK(tracker.GetResIDFromAddr(0x1250000 - 1) == none); + CHECK(tracker.GetResIDFromAddr(0x1250000) == make_idoffs(b, 0ULL)); + CHECK(tracker.GetResIDFromAddr(0x1250001) == make_idoffs(b, 1ULL)); + CHECK(tracker.GetResIDFromAddr(0x1250000 + 127) == make_idoffs(b, 127ULL)); + CHECK(tracker.GetResIDFromAddr(0x1250000 + 128) == none); + + tracker.RemoveFrom(MakeRange(b, 0x1250000, 128)); + + CHECK(tracker.GetResIDFromAddr(0x1230000 - 1) == none); + CHECK(tracker.GetResIDFromAddr(0x1230000) == make_idoffs(a, 0ULL)); + CHECK(tracker.GetResIDFromAddr(0x1230001) == make_idoffs(a, 1ULL)); + CHECK(tracker.GetResIDFromAddr(0x1230000 + 127) == make_idoffs(a, 127ULL)); + CHECK(tracker.GetResIDFromAddr(0x1230000 + 128) == none); + + CHECK(tracker.GetResIDFromAddr(0x1250000 - 1) == none); + CHECK(tracker.GetResIDFromAddr(0x1250000) == none); + CHECK(tracker.GetResIDFromAddr(0x1250001) == none); + CHECK(tracker.GetResIDFromAddr(0x1250000 + 127) == none); + CHECK(tracker.GetResIDFromAddr(0x1250000 + 128) == none); + + tracker.AddTo(MakeRange(c, 0x1270000, 128)); + + CHECK(tracker.GetResIDFromAddr(0x1230000 - 1) == none); + CHECK(tracker.GetResIDFromAddr(0x1230000) == make_idoffs(a, 0ULL)); + CHECK(tracker.GetResIDFromAddr(0x1230001) == make_idoffs(a, 1ULL)); + CHECK(tracker.GetResIDFromAddr(0x1230000 + 127) == make_idoffs(a, 127ULL)); + CHECK(tracker.GetResIDFromAddr(0x1230000 + 128) == none); + + CHECK(tracker.GetResIDFromAddr(0x1250000 - 1) == none); + CHECK(tracker.GetResIDFromAddr(0x1250000) == none); + CHECK(tracker.GetResIDFromAddr(0x1250001) == none); + CHECK(tracker.GetResIDFromAddr(0x1250000 + 127) == none); + CHECK(tracker.GetResIDFromAddr(0x1250000 + 128) == none); + + CHECK(tracker.GetResIDFromAddr(0x1270000 - 1) == none); + CHECK(tracker.GetResIDFromAddr(0x1270000) == make_idoffs(c, 0ULL)); + CHECK(tracker.GetResIDFromAddr(0x1270001) == make_idoffs(c, 1ULL)); + CHECK(tracker.GetResIDFromAddr(0x1270000 + 127) == make_idoffs(c, 127ULL)); + CHECK(tracker.GetResIDFromAddr(0x1270000 + 128) == none); + + EXPECT_ERROR(); + + // wrong ID, don't remove + tracker.RemoveFrom(MakeRange(g, 0x1270000, 128)); + + CHECK(tracker.GetResIDFromAddr(0x1230000 - 1) == none); + CHECK(tracker.GetResIDFromAddr(0x1230000) == make_idoffs(a, 0ULL)); + CHECK(tracker.GetResIDFromAddr(0x1230001) == make_idoffs(a, 1ULL)); + CHECK(tracker.GetResIDFromAddr(0x1230000 + 127) == make_idoffs(a, 127ULL)); + CHECK(tracker.GetResIDFromAddr(0x1230000 + 128) == none); + + CHECK(tracker.GetResIDFromAddr(0x1250000 - 1) == none); + CHECK(tracker.GetResIDFromAddr(0x1250000) == none); + CHECK(tracker.GetResIDFromAddr(0x1250001) == none); + CHECK(tracker.GetResIDFromAddr(0x1250000 + 127) == none); + CHECK(tracker.GetResIDFromAddr(0x1250000 + 128) == none); + + CHECK(tracker.GetResIDFromAddr(0x1270000 - 1) == none); + CHECK(tracker.GetResIDFromAddr(0x1270000) == make_idoffs(c, 0ULL)); + CHECK(tracker.GetResIDFromAddr(0x1270001) == make_idoffs(c, 1ULL)); + CHECK(tracker.GetResIDFromAddr(0x1270000 + 127) == make_idoffs(c, 127ULL)); + CHECK(tracker.GetResIDFromAddr(0x1270000 + 128) == none); + + EXPECT_ERROR(); + + // wrong address, don't remove + tracker.RemoveFrom(MakeRange(a, 0x1000, 128)); + + CHECK(tracker.GetResIDFromAddr(0x1230000 - 1) == none); + CHECK(tracker.GetResIDFromAddr(0x1230000) == make_idoffs(a, 0ULL)); + CHECK(tracker.GetResIDFromAddr(0x1230001) == make_idoffs(a, 1ULL)); + CHECK(tracker.GetResIDFromAddr(0x1230000 + 127) == make_idoffs(a, 127ULL)); + CHECK(tracker.GetResIDFromAddr(0x1230000 + 128) == none); + + CHECK(tracker.GetResIDFromAddr(0x1250000 - 1) == none); + CHECK(tracker.GetResIDFromAddr(0x1250000) == none); + CHECK(tracker.GetResIDFromAddr(0x1250001) == none); + CHECK(tracker.GetResIDFromAddr(0x1250000 + 127) == none); + CHECK(tracker.GetResIDFromAddr(0x1250000 + 128) == none); + + CHECK(tracker.GetResIDFromAddr(0x1270000 - 1) == none); + CHECK(tracker.GetResIDFromAddr(0x1270000) == make_idoffs(c, 0ULL)); + CHECK(tracker.GetResIDFromAddr(0x1270001) == make_idoffs(c, 1ULL)); + CHECK(tracker.GetResIDFromAddr(0x1270000 + 127) == make_idoffs(c, 127ULL)); + CHECK(tracker.GetResIDFromAddr(0x1270000 + 128) == none); + } + + SECTION("Insertion order doesn't affect return value") + { + // smallest-to-largest + tracker.AddTo(MakeRange(a, 0x1230000, 128)); + tracker.AddTo(MakeRange(b, 0x1230000, 256)); + tracker.AddTo(MakeRange(c, 0x1230000, 512)); + + CHECK(tracker.GetResIDFromAddr(0x1230000) == make_idoffs(c, 0ULL)); + + tracker.Clear(); + + // largest-to-smallest + tracker.AddTo(MakeRange(c, 0x1230000, 512)); + tracker.AddTo(MakeRange(b, 0x1230000, 256)); + tracker.AddTo(MakeRange(a, 0x1230000, 128)); + + CHECK(tracker.GetResIDFromAddr(0x1230000) == make_idoffs(c, 0ULL)); + + tracker.Clear(); + + // out-of-order, largest last + tracker.AddTo(MakeRange(b, 0x1230000, 256)); + tracker.AddTo(MakeRange(a, 0x1230000, 128)); + tracker.AddTo(MakeRange(c, 0x1230000, 512)); + + CHECK(tracker.GetResIDFromAddr(0x1230000) == make_idoffs(c, 0ULL)); + + tracker.Clear(); + + // out-of-order, smallest last + tracker.AddTo(MakeRange(b, 0x1230000, 256)); + tracker.AddTo(MakeRange(c, 0x1230000, 512)); + tracker.AddTo(MakeRange(a, 0x1230000, 128)); + + CHECK(tracker.GetResIDFromAddr(0x1230000) == make_idoffs(c, 0ULL)); + + tracker.Clear(); + + // with a pre-existing address before the ranges + tracker.AddTo(MakeRange(d, 0x1200000, 512)); + tracker.AddTo(MakeRange(c, 0x1230000, 512)); + tracker.AddTo(MakeRange(b, 0x1230000, 256)); + tracker.AddTo(MakeRange(a, 0x1230000, 128)); + + CHECK(tracker.GetResIDFromAddr(0x1230000) == make_idoffs(c, 0ULL)); + + tracker.Clear(); + + // with a pre-existing address after the ranges + tracker.AddTo(MakeRange(d, 0x1250000, 512)); + tracker.AddTo(MakeRange(c, 0x1230000, 512)); + tracker.AddTo(MakeRange(b, 0x1230000, 256)); + tracker.AddTo(MakeRange(a, 0x1230000, 128)); + + CHECK(tracker.GetResIDFromAddr(0x1230000) == make_idoffs(c, 0ULL)); + + tracker.Clear(); + } + + SECTION("OOB") + { + tracker.AddTo(MakeRange(a, 0x1230000, 128, 128)); + + CHECK(tracker.GetResIDFromAddrAllowOutOfBounds(0x1230001) == make_idoffs(a, 1ULL)); + CHECK(tracker.GetResIDFromAddrAllowOutOfBounds(0x1230000 + 127) == make_idoffs(a, 127ULL)); + CHECK(tracker.GetResIDFromAddrAllowOutOfBounds(0x1230000 + 128) == make_idoffs(a, 128ULL)); + CHECK(tracker.GetResIDFromAddrAllowOutOfBounds(0x1230000 + 255) == make_idoffs(a, 255ULL)); + CHECK(tracker.GetResIDFromAddrAllowOutOfBounds(0x1230000 + 256) == none); + + tracker.RemoveFrom(MakeRange(a, 0x1230000, 128, 128)); + + CHECK(tracker.GetResIDFromAddrAllowOutOfBounds(0x1230001) == none); + CHECK(tracker.GetResIDFromAddrAllowOutOfBounds(0x1230000 + 127) == none); + CHECK(tracker.GetResIDFromAddrAllowOutOfBounds(0x1230000 + 128) == none); + CHECK(tracker.GetResIDFromAddrAllowOutOfBounds(0x1230000 + 255) == none); + CHECK(tracker.GetResIDFromAddrAllowOutOfBounds(0x1230000 + 256) == none); + + tracker.AddTo(MakeRange(a, 0x1230000, 0x10000, 0x10000)); + tracker.AddTo(MakeRange(b, 0x1250000, 0x10000, 0x10000)); + + CHECK(tracker.GetResIDFromAddrAllowOutOfBounds(0x1230001) == make_idoffs(a, 0x1ULL)); + CHECK(tracker.GetResIDFromAddrAllowOutOfBounds(0x1240000) == make_idoffs(a, 0x10000ULL)); + CHECK(tracker.GetResIDFromAddrAllowOutOfBounds(0x1240001) == make_idoffs(a, 0x10001ULL)); + CHECK(tracker.GetResIDFromAddrAllowOutOfBounds(0x124ffff) == make_idoffs(a, 0x1ffffULL)); + CHECK(tracker.GetResIDFromAddrAllowOutOfBounds(0x1250000) == make_idoffs(b, 0ULL)); + } + + SECTION("co-sited overlap returning largest") + { + auto checker = [&tracker, none](ResourceId id) { + CHECK(tracker.GetResIDFromAddr(0x1230000 - 1) == none); + CHECK(tracker.GetResIDFromAddr(0x1230000) == make_idoffs(id, 0ULL)); + CHECK(tracker.GetResIDFromAddr(0x1230001) == make_idoffs(id, 1ULL)); + CHECK(tracker.GetResIDFromAddr(0x1230010) == make_idoffs(id, 0x10ULL)); + CHECK(tracker.GetResIDFromAddr(0x1230000 + 127) == make_idoffs(id, 127ULL)); + + // check the range of a we expect + if(id == a) + { + CHECK(tracker.GetResIDFromAddr(0x1230000 + 128) == make_idoffs(id, 128ULL)); + CHECK(tracker.GetResIDFromAddr(0x1230000 + 255) == make_idoffs(id, 255ULL)); + } + else + { + CHECK(tracker.GetResIDFromAddr(0x1230000 + 128) == none); + CHECK(tracker.GetResIDFromAddr(0x1230000 + 255) == none); + } + }; + + SECTION("big before small") + { + tracker.AddTo(MakeRange(a, 0x1230000, 256)); + tracker.AddTo(MakeRange(b, 0x1230000, 128)); + + // should find a regardless of added order + checker(a); + + SECTION("remove a") + { + // if a is removed, we now find b + tracker.RemoveFrom(MakeRange(a, 0x1230000, 256)); + + checker(b); + } + + SECTION("remove b") + { + // if b is removed, we still find a + tracker.RemoveFrom(MakeRange(b, 0x1230000, 128)); + + checker(a); + } + } + + SECTION("small before big") + { + tracker.AddTo(MakeRange(b, 0x1230000, 128)); + tracker.AddTo(MakeRange(a, 0x1230000, 256)); + + // should find a regardless of added order + checker(a); + + SECTION("remove a") + { + // if a is removed, we now find b + tracker.RemoveFrom(MakeRange(a, 0x1230000, 256)); + + checker(b); + } + + SECTION("remove b") + { + // if b is removed, we still find a + tracker.RemoveFrom(MakeRange(b, 0x1230000, 128)); + + checker(a); + } + } + } + + SECTION("Partially overlaping ranges that aren't super/subset") + { + tracker.AddTo(MakeRange(c, 0x12000, 0x0800)); + tracker.AddTo(MakeRange(d, 0x12600, 0x0800)); + tracker.AddTo(MakeRange(e, 0x12800, 0x0200)); + + CHECK(tracker.GetResIDFromAddr(0x12000) == make_idoffs(c, 0ULL)); + CHECK(tracker.GetResIDFromAddr(0x12100) == make_idoffs(c, 0x100ULL)); + CHECK(tracker.GetResIDFromAddr(0x125ff) == make_idoffs(c, 0x5ffULL)); + CHECK(tracker.GetResIDFromAddr(0x12600) == make_idoffs(d, 0ULL)); + CHECK(tracker.GetResIDFromAddr(0x12700) == make_idoffs(d, 0x100ULL)); + CHECK(tracker.GetResIDFromAddr(0x127ff) == make_idoffs(d, 0x1ffULL)); + CHECK(tracker.GetResIDFromAddr(0x12800) == make_idoffs(e, 0ULL)); + CHECK(tracker.GetResIDFromAddr(0x12900) == make_idoffs(e, 0x100ULL)); + CHECK(tracker.GetResIDFromAddr(0x129ff) == make_idoffs(e, 0x1ffULL)); + } + + SECTION("lots of overlap and removals") + { + tracker.AddTo(MakeRange(a, 0x12300000, 100)); + tracker.AddTo(MakeRange(b, 0x12300000, 200)); + tracker.AddTo(MakeRange(c, 0x12300000, 300)); + tracker.AddTo(MakeRange(d, 0x12300000, 400)); + tracker.AddTo(MakeRange(e, 0x12300000, 500)); + tracker.AddTo(MakeRange(f, 0x12300000, 600)); + + CHECK(tracker.GetResIDFromAddr(0x12300000 - 1) == none); + CHECK(tracker.GetResIDFromAddr(0x12300000) == make_idoffs(f, 0ULL)); + CHECK(tracker.GetResIDFromAddr(0x12300f00) == none); + + tracker.RemoveFrom(MakeRange(c, 0x12300000, 300)); + + CHECK(tracker.GetResIDFromAddr(0x12300000 - 1) == none); + CHECK(tracker.GetResIDFromAddr(0x12300000) == make_idoffs(f, 0ULL)); + CHECK(tracker.GetResIDFromAddr(0x12300f00) == none); + + tracker.RemoveFrom(MakeRange(f, 0x12300000, 600)); + + CHECK(tracker.GetResIDFromAddr(0x12300000 - 1) == none); + CHECK(tracker.GetResIDFromAddr(0x12300000) == make_idoffs(e, 0ULL)); + CHECK(tracker.GetResIDFromAddr(0x12300f00) == none); + + tracker.RemoveFrom(MakeRange(a, 0x12300000, 100)); + + CHECK(tracker.GetResIDFromAddr(0x12300000 - 1) == none); + CHECK(tracker.GetResIDFromAddr(0x12300000) == make_idoffs(e, 0ULL)); + CHECK(tracker.GetResIDFromAddr(0x12300f00) == none); + + tracker.RemoveFrom(MakeRange(d, 0x12300000, 100)); + + CHECK(tracker.GetResIDFromAddr(0x12300000 - 1) == none); + CHECK(tracker.GetResIDFromAddr(0x12300000) == make_idoffs(e, 0ULL)); + CHECK(tracker.GetResIDFromAddr(0x12300f00) == none); + + tracker.RemoveFrom(MakeRange(e, 0x12300000, 100)); + + CHECK(tracker.GetResIDFromAddr(0x12300000 - 1) == none); + CHECK(tracker.GetResIDFromAddr(0x12300000) == make_idoffs(b, 0ULL)); + CHECK(tracker.GetResIDFromAddr(0x12300f00) == none); + + tracker.RemoveFrom(MakeRange(b, 0x12300000, 200)); + + CHECK(tracker.GetResIDFromAddr(0x12300000 - 1) == none); + CHECK(tracker.GetResIDFromAddr(0x12300000) == none); + CHECK(tracker.GetResIDFromAddr(0x12300f00) == none); + } +} + +#endif diff --git a/renderdoc/core/gpu_address_range_tracker.h b/renderdoc/core/gpu_address_range_tracker.h index 096e946ba..d5e410458 100644 --- a/renderdoc/core/gpu_address_range_tracker.h +++ b/renderdoc/core/gpu_address_range_tracker.h @@ -36,13 +36,7 @@ struct GPUAddressRange Address start, realEnd, oobEnd; ResourceId id; - bool operator<(const Address &o) const - { - if(o < start) - return true; - - return false; - } + bool operator<(const Address &o) const { return (start < o); } }; struct GPUAddressRangeTracker @@ -52,15 +46,43 @@ struct GPUAddressRangeTracker GPUAddressRangeTracker(const GPUAddressRangeTracker &) = delete; GPUAddressRangeTracker &operator=(const GPUAddressRangeTracker &) = delete; - rdcarray addresses; - Threading::RWLock addressLock; - void AddTo(const GPUAddressRange &range); void RemoveFrom(const GPUAddressRange &range); - void GetResIDFromAddr(GPUAddressRange::Address addr, ResourceId &id, uint64_t &offs); - void GetResIDFromAddrAllowOutOfBounds(GPUAddressRange::Address addr, ResourceId &id, - uint64_t &offs); + void Clear(); + rdcarray GetAddresses(); + rdcarray GetIDs(); + + void GetResIDFromAddr(GPUAddressRange::Address addr, ResourceId &id, uint64_t &offs) + { + return GetResIDFromAddr(addr, id, offs); + } + void GetResIDFromAddrAllowOutOfBounds(GPUAddressRange::Address addr, ResourceId &id, uint64_t &offs) + { + return GetResIDFromAddr(addr, id, offs); + } + + rdcpair GetResIDFromAddr(GPUAddressRange::Address addr) + { + rdcpair ret; + GetResIDFromAddr(addr, ret.first, ret.second); + return ret; + } + rdcpair GetResIDFromAddrAllowOutOfBounds(GPUAddressRange::Address addr) + { + rdcpair ret; + GetResIDFromAddrAllowOutOfBounds(addr, ret.first, ret.second); + return ret; + } void GetResIDBoundForAddr(GPUAddressRange::Address addr, ResourceId &lower, GPUAddressRange::Address &lowerVA, ResourceId &upper, GPUAddressRange::Address &upperVA); + +private: + rdcarray addresses; + Threading::RWLock addressLock; + + template + void GetResIDFromAddr(GPUAddressRange::Address addr, ResourceId &id, uint64_t &offs); + + size_t FindLastRangeBeforeOrAtAddress(GPUAddressRange::Address start); }; diff --git a/renderdoc/driver/d3d12/d3d12_debug.cpp b/renderdoc/driver/d3d12/d3d12_debug.cpp index 42e6b2e91..7029e9e25 100644 --- a/renderdoc/driver/d3d12/d3d12_debug.cpp +++ b/renderdoc/driver/d3d12/d3d12_debug.cpp @@ -1988,7 +1988,7 @@ D3D12_CPU_DESCRIPTOR_HANDLE D3D12DebugManager::GetUAVClearHandle(CBVUAVSRVSlot s return ret; } -void D3D12DebugManager::PrepareExecuteIndirectPatching(const GPUAddressRangeTracker &origAddresses) +void D3D12DebugManager::PrepareExecuteIndirectPatching(GPUAddressRangeTracker &origAddresses) { D3D12ShaderCache *shaderCache = m_pDevice->GetShaderCache(); @@ -2052,7 +2052,7 @@ void D3D12DebugManager::PrepareExecuteIndirectPatching(const GPUAddressRangeTrac }; rdcarray buffers; - for(const GPUAddressRange &addr : origAddresses.addresses) + for(const GPUAddressRange &addr : origAddresses.GetAddresses()) { buffermapping b = {}; b.origBase = addr.start; diff --git a/renderdoc/driver/d3d12/d3d12_debug.h b/renderdoc/driver/d3d12/d3d12_debug.h index 2cf3b7385..0496abdf5 100644 --- a/renderdoc/driver/d3d12/d3d12_debug.h +++ b/renderdoc/driver/d3d12/d3d12_debug.h @@ -200,7 +200,7 @@ public: void PrepareTextureSampling(ID3D12Resource *resource, CompType typeCast, int &resType, BarrierSet &barrierSet); - void PrepareExecuteIndirectPatching(const GPUAddressRangeTracker &origAddresses); + void PrepareExecuteIndirectPatching(GPUAddressRangeTracker &origAddresses); MeshDisplayPipelines CacheMeshDisplayPipelines(const MeshFormat &primary, const MeshFormat &secondary); diff --git a/renderdoc/driver/d3d12/d3d12_device.cpp b/renderdoc/driver/d3d12/d3d12_device.cpp index 570686260..5d713e82b 100644 --- a/renderdoc/driver/d3d12/d3d12_device.cpp +++ b/renderdoc/driver/d3d12/d3d12_device.cpp @@ -3259,9 +3259,8 @@ void WrappedID3D12Device::UploadBLASBufferAddresses() rdcarray blasAddressPair; D3D12ResourceManager *resManager = GetResourceManager(); - for(size_t i = 0; i < m_OrigGPUAddresses.addresses.size(); i++) + for(GPUAddressRange addressRange : m_OrigGPUAddresses.GetAddresses()) { - GPUAddressRange addressRange = m_OrigGPUAddresses.addresses[i]; ResourceId resId = addressRange.id; if(resManager->HasLiveResource(resId)) { diff --git a/renderdoc/driver/d3d12/d3d12_manager.cpp b/renderdoc/driver/d3d12/d3d12_manager.cpp index c66dc6f49..a6711ea44 100644 --- a/renderdoc/driver/d3d12/d3d12_manager.cpp +++ b/renderdoc/driver/d3d12/d3d12_manager.cpp @@ -2197,7 +2197,7 @@ PatchedRayDispatch D3D12RTManager::PatchIndirectRayDispatch( return ret; } -void D3D12RTManager::PrepareRayDispatchBuffer(const GPUAddressRangeTracker *origAddresses) +void D3D12RTManager::PrepareRayDispatchBuffer(GPUAddressRangeTracker *origAddresses) { SCOPED_LOCK(m_LookupBufferLock); if(m_LookupBufferDirty || origAddresses) @@ -2230,10 +2230,13 @@ void D3D12RTManager::PrepareRayDispatchBuffer(const GPUAddressRangeTracker *orig const size_t RootSigOffset = lookupData.size(); lookupData.resize(lookupData.size() + m_UniqueLocalRootSigs.size() * RootSigStride); + rdcarray addresses; + const size_t PatchAddrOffset = lookupData.size(); if(origAddresses) { - lookupData.resize(lookupData.size() + sizeof(BlasAddressPair) * origAddresses->addresses.size()); + addresses = origAddresses->GetAddresses(); + lookupData.resize(lookupData.size() + sizeof(BlasAddressPair) * addresses.size()); } else { @@ -2264,9 +2267,9 @@ void D3D12RTManager::PrepareRayDispatchBuffer(const GPUAddressRangeTracker *orig m_NumPatchingAddrs = 0; - for(size_t i = 0; origAddresses && i < origAddresses->addresses.size(); i++) + for(size_t i = 0; origAddresses && i < addresses.size(); i++) { - GPUAddressRange addressRange = origAddresses->addresses[i]; + GPUAddressRange addressRange = addresses[i]; ResourceId resId = addressRange.id; if(m_wrappedDevice->GetResourceManager()->HasLiveResource(resId)) { diff --git a/renderdoc/driver/d3d12/d3d12_manager.h b/renderdoc/driver/d3d12/d3d12_manager.h index e7c44492b..79f183103 100644 --- a/renderdoc/driver/d3d12/d3d12_manager.h +++ b/renderdoc/driver/d3d12/d3d12_manager.h @@ -1241,7 +1241,7 @@ public: void RegisterExportDatabase(D3D12ShaderExportDatabase *db); void UnregisterExportDatabase(D3D12ShaderExportDatabase *db); - void PrepareRayDispatchBuffer(const GPUAddressRangeTracker *origAddresses); + void PrepareRayDispatchBuffer(GPUAddressRangeTracker *origAddresses); ASBuildData *CopyBuildInputs(ID3D12GraphicsCommandList4 *unwrappedCmd, const D3D12_BUILD_RAYTRACING_ACCELERATION_STRUCTURE_INPUTS &inputs); diff --git a/renderdoc/driver/d3d12/d3d12_resources.cpp b/renderdoc/driver/d3d12/d3d12_resources.cpp index 7e68be058..57373bbfc 100644 --- a/renderdoc/driver/d3d12/d3d12_resources.cpp +++ b/renderdoc/driver/d3d12/d3d12_resources.cpp @@ -447,22 +447,19 @@ bool WrappedID3D12Resource::DeleteAccStructAtOffset(D3D12BufferOffset bufferOffs void WrappedID3D12Resource::RefBuffers(D3D12ResourceManager *rm) { // only buffers go into m_Addresses - SCOPED_READLOCK(m_Addresses.addressLock); - for(size_t i = 0; i < m_Addresses.addresses.size(); i++) - rm->MarkResourceFrameReferenced(m_Addresses.addresses[i].id, eFrameRef_Read); + for(ResourceId id : m_Addresses.GetIDs()) + rm->MarkResourceFrameReferenced(id, eFrameRef_Read); } void WrappedID3D12Resource::GetMappableIDs(D3D12ResourceManager *rm, const std::unordered_set &refdIDs, std::unordered_set &mappableIDs) { - SCOPED_READLOCK(m_Addresses.addressLock); - for(size_t i = 0; i < m_Addresses.addresses.size(); i++) + for(ResourceId id : m_Addresses.GetIDs()) { - if(refdIDs.find(m_Addresses.addresses[i].id) != refdIDs.end()) + if(refdIDs.find(id) != refdIDs.end()) { - WrappedID3D12Resource *resource = - (WrappedID3D12Resource *)rm->GetCurrentResource(m_Addresses.addresses[i].id); + WrappedID3D12Resource *resource = (WrappedID3D12Resource *)rm->GetCurrentResource(id); mappableIDs.insert(resource->GetMappableID()); } } @@ -472,11 +469,7 @@ rdcarray WrappedID3D12Resource::AddRefBuffersBeforeCapture(D3D { rdcarray ret; - rdcarray addresses; - { - SCOPED_READLOCK(m_Addresses.addressLock); - addresses = m_Addresses.addresses; - } + rdcarray addresses = m_Addresses.GetAddresses(); for(size_t i = 0; i < addresses.size(); i++) {