File size: 5,020 Bytes
4a28d4d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 | // Copyright (c) OpenMMLab. All rights reserved.
#pragma once
#include "src/turbomind/kernels/core/common.h"
#include "src/turbomind/kernels/core/math.h"
#include <iostream>
namespace turbomind {
template<int C, int S, int AccessC, int WarpCount>
struct ThreadMapQ {
static constexpr int kWarpCount = WarpCount;
static constexpr int kAccessC = AccessC;
static constexpr int kWarpThreadC = C / kAccessC;
static constexpr int kWarpThreadS = WARP_SIZE / kWarpThreadC;
static_assert(kWarpThreadC <= WARP_SIZE);
static constexpr int kWarpAccessC = kWarpThreadC * kAccessC; // C
static constexpr int kWarpAccessS = kWarpThreadS;
static constexpr int kWarpIterC = C / kWarpAccessC; // 1
static constexpr int kWarpIterS = S / kWarpAccessS;
static constexpr int kWarpC = 1;
static constexpr int kWarpS = kWarpCount;
static constexpr int kIterC = kWarpIterC / kWarpC; // 1
static constexpr int kIterS = std::max(kWarpIterS / kWarpS, 1);
static constexpr int kFootprintC = kWarpAccessC * kIterC; // C
static constexpr int kFootprintS = kWarpAccessS * kIterS;
static constexpr int kDeltaC = kWarpAccessC;
static constexpr int kDeltaS = kWarpAccessS;
__device__ static int2 get_offset(int warp_id, int lane_id)
{
int warp_offset_c = warp_id % kWarpC;
int warp_offset_s = warp_id / kWarpC;
int warp_thread_offset_c = lane_id % kWarpThreadC;
int warp_thread_offset_s = lane_id / kWarpThreadC;
int cta_thread_offset_c = kFootprintC * warp_offset_c + warp_thread_offset_c * kAccessC;
int cta_thread_offset_s = kFootprintS * warp_offset_s + warp_thread_offset_s;
return {cta_thread_offset_c, cta_thread_offset_s};
}
};
template<int DimC, int DimS, int AccessC, int WarpCount, int WarpThreadC = lowbit(DimC) / AccessC, int WarpC = 1>
struct RakedThreadMap {
static constexpr int kDimC = DimC;
static constexpr int kDimS = DimS;
static constexpr int kWarpCount = WarpCount;
static constexpr int kAccessC = AccessC;
static constexpr int kWarpThreadC = WarpThreadC;
static constexpr int kWarpThreadS = WARP_SIZE / kWarpThreadC;
static_assert(WARP_SIZE % kWarpThreadC == 0);
static constexpr int kWarpAccessC = kWarpThreadC * kAccessC;
static constexpr int kWarpAccessS = kWarpThreadS;
static constexpr int kWarpIterC = cdiv(kDimC, kWarpAccessC);
static constexpr int kWarpIterS = cdiv(kDimS, kWarpAccessS);
static constexpr int kWarpC = WarpC;
static constexpr int kWarpS = kWarpCount / kWarpC;
static_assert(kWarpCount % kWarpC == 0);
static constexpr int kIterC = cdiv(kWarpIterC, kWarpC);
static constexpr int kIterS = cdiv(kWarpIterS, kWarpS);
// Allow partial tile when there is ONLY 1 iteration
static_assert(kDimC % kWarpAccessC == 0 || kIterC == 1);
static constexpr bool kPartialC = kDimC % kWarpAccessC != 0;
static constexpr int kFootprintC = kWarpAccessC * kIterC;
static constexpr int kFootprintS = kWarpAccessS * kIterS;
static constexpr int kDeltaC = kWarpAccessC;
static constexpr int kDeltaS = kWarpAccessS;
// static constexpr int kDeltaC = kWarpAccessC * kWarpC;
// static constexpr int kDeltaS = kWarpAccessS * kWarpS;
__device__ static int2 get_offset(int warp_id, int lane_id)
{
int warp_offset_c = warp_id % kWarpC;
int warp_offset_s = warp_id / kWarpC;
int warp_thread_offset_c = lane_id % kWarpThreadC;
int warp_thread_offset_s = lane_id / kWarpThreadC;
int cta_thread_offset_c = kFootprintC * warp_offset_c + warp_thread_offset_c * kAccessC;
int cta_thread_offset_s = kFootprintS * warp_offset_s + warp_thread_offset_s;
// int cta_thread_offset_c = kWarpAccessC * warp_offset_c + warp_thread_offset_c * kAccessC;
// int cta_thread_offset_s = kWarpAccessS * warp_offset_s + warp_thread_offset_s;
return {cta_thread_offset_c, cta_thread_offset_s};
}
};
namespace {
template<class TMap>
void Print(TMap)
{
std::cout << " warps: " << TMap::kWarpCount << "\n";
std::cout << " shape: (" << TMap::kDimC << ", " << TMap::kDimS << ")\n";
std::cout << " access: (" << TMap::kAccessC << ", " << 1 << ")\n";
std::cout << "warpThread: (" << TMap::kWarpThreadC << ", " << TMap::kWarpThreadS << ")\n";
std::cout << "warpAccess: (" << TMap::kWarpAccessC << ", " << TMap::kWarpAccessS << ")\n";
std::cout << " warpIter: (" << TMap::kWarpIterC << ", " << TMap::kWarpIterS << ")\n";
std::cout << " warp: (" << TMap::kWarpC << ", " << TMap::kWarpS << ")\n";
std::cout << " iter: (" << TMap::kIterC << ", " << TMap::kIterS << ")\n";
std::cout << " footprint: (" << TMap::kFootprintC << ", " << TMap::kFootprintS << ")\n";
std::cout << " delta: (" << TMap::kDeltaC << ", " << TMap::kDeltaS << ")\n";
std::cout << " partialC: " << TMap::kPartialC << "\n";
}
} // namespace
} // namespace turbomind
|