refactor: root kernel includes at csrc/kernels
- add the kernels directory to CMake target include paths and drop all ../-relative includes in kernel sources - reference shared primitives as common/*.cuh and the fp8 type header as fp8/common.h - update standalone test nvcc commands in file headers and cuda_kernels.md to -I csrc/kernels
This commit is contained in:
@@ -2,8 +2,8 @@
|
||||
#include <cuda_bf16.h>
|
||||
#include <float.h>
|
||||
#include "common.h"
|
||||
#include "common/reduce.cuh"
|
||||
#include "layout_policies.cuh"
|
||||
#include "../common/reduce.cuh"
|
||||
|
||||
namespace astrai {
|
||||
namespace attention {
|
||||
|
||||
@@ -3,8 +3,8 @@
|
||||
#include <cuda_fp16.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#include "../common/cp_async.cuh"
|
||||
#include "../common/mma.cuh"
|
||||
#include "common/cp_async.cuh"
|
||||
#include "common/mma.cuh"
|
||||
|
||||
// Predicated cp.async (4-operand form) requires CUDA 11.2+.
|
||||
// bf16 mma.sync requires sm_80+ (guarded at build time by ASTRAI_NO_MMA).
|
||||
|
||||
@@ -2,8 +2,8 @@
|
||||
#include <cfloat>
|
||||
#include <cuda_bf16.h>
|
||||
#include "common.h"
|
||||
#include "common/reduce.cuh"
|
||||
#include "layout_policies.cuh"
|
||||
#include "../common/reduce.cuh"
|
||||
|
||||
namespace astrai {
|
||||
namespace attention {
|
||||
|
||||
@@ -11,8 +11,8 @@
|
||||
#include <cuda_runtime.h>
|
||||
#include <type_traits>
|
||||
|
||||
#include "../common/cp_async.cuh"
|
||||
#include "common.h"
|
||||
#include "common/cp_async.cuh"
|
||||
#include "gemm/epilogue.cuh"
|
||||
#include "gemm/load.cuh"
|
||||
#include "gemm/mainloop.cuh"
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
// Collective epilogue: fused bias, the bf16 scatter of the fp32 accumulators
|
||||
// through the reclaimed operand shared memory, and the coalesced copy-out.
|
||||
|
||||
#include "../common.h"
|
||||
#include "fp8/common.h"
|
||||
#include "policy.cuh"
|
||||
|
||||
namespace astrai {
|
||||
|
||||
@@ -5,8 +5,8 @@
|
||||
// The staging invariants and the swizzle derivation live in
|
||||
// docs/developer/cuda_kernels.md.
|
||||
|
||||
#include "../../common/cp_async.cuh"
|
||||
#include "../common.h"
|
||||
#include "common/cp_async.cuh"
|
||||
#include "fp8/common.h"
|
||||
#include "policy.cuh"
|
||||
|
||||
namespace astrai {
|
||||
|
||||
@@ -7,8 +7,8 @@
|
||||
|
||||
#include <type_traits>
|
||||
|
||||
#include "../../common/mma.cuh"
|
||||
#include "../common.h"
|
||||
#include "common/mma.cuh"
|
||||
#include "fp8/common.h"
|
||||
#include "load.cuh"
|
||||
#include "policy.cuh"
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@
|
||||
|
||||
#include <type_traits>
|
||||
|
||||
#include "../common.h"
|
||||
#include "fp8/common.h"
|
||||
|
||||
namespace astrai {
|
||||
namespace fp8 {
|
||||
|
||||
@@ -8,7 +8,7 @@
|
||||
#include <mutex>
|
||||
#include <unordered_map>
|
||||
|
||||
#include "../common/device.cuh"
|
||||
#include "common/device.cuh"
|
||||
#include "gemm.cuh"
|
||||
#include "quantize.cuh"
|
||||
|
||||
|
||||
@@ -10,7 +10,7 @@
|
||||
#include <cstdint>
|
||||
|
||||
#include "common.h"
|
||||
#include "../common/reduce.cuh"
|
||||
#include "common/reduce.cuh"
|
||||
|
||||
namespace astrai {
|
||||
namespace fp8 {
|
||||
|
||||
Reference in New Issue
Block a user