mirror of
https://github.com/gentoo-mirror/gentoo.git
synced 2026-09-24 04:59:14 -07:00
sci-ml/caffe2: Fix use of undeclared identifier 'CHECK_NOSPARSE_CONTIGUOUS_CUDA' with USE='-flash'
Bug: https://github.com/pytorch/pytorch/issues/160826 Signed-off-by: Sv. Lockal <lockalsash@gmail.com> Part-of: https://github.com/gentoo/gentoo/pull/43468 Closes: https://github.com/gentoo/gentoo/pull/43468 Signed-off-by: Alfredo Tupone <tupone@gentoo.org>
This commit is contained in:
committed by
Alfredo Tupone
parent
36b0dca1ca
commit
906973431e
@@ -147,6 +147,7 @@ PATCHES=(
|
||||
"${FILESDIR}"/${P}-cmake.patch
|
||||
"${FILESDIR}"/${PN}-2.7.0-glog-0.7.1.patch
|
||||
"${FILESDIR}"/${PN}-2.7.1-aotriton-fixes.patch
|
||||
"${FILESDIR}"/${PN}-2.8.0-rocm-minus-flash.patch
|
||||
)
|
||||
|
||||
src_prepare() {
|
||||
|
||||
86
sci-ml/caffe2/files/caffe2-2.8.0-rocm-minus-flash.patch
Normal file
86
sci-ml/caffe2/files/caffe2-2.8.0-rocm-minus-flash.patch
Normal file
@@ -0,0 +1,86 @@
|
||||
Fix use of undeclared identifier 'CHECK_NOSPARSE_CONTIGUOUS_CUDA' with USE='-flash'
|
||||
|
||||
Bug: https://github.com/pytorch/pytorch/issues/160826
|
||||
--- a/aten/src/ATen/native/transformers/cuda/attention.cu
|
||||
+++ b/aten/src/ATen/native/transformers/cuda/attention.cu
|
||||
@@ -71,6 +71,7 @@
|
||||
#include <ATen/native/transformers/cuda/sdp_utils.h>
|
||||
#include <ATen/native/transformers/sdp_utils_cpp.h>
|
||||
|
||||
+#include <ATen/native/transformers/flash_api_common.h>
|
||||
#ifdef USE_FLASH_ATTENTION
|
||||
// FlashAttention Specific Imports
|
||||
#include <ATen/native/transformers/cuda/flash_attn/flash_api.h>
|
||||
--- a/aten/src/ATen/native/transformers/cuda/attention_backward.cu
|
||||
+++ b/aten/src/ATen/native/transformers/cuda/attention_backward.cu
|
||||
@@ -33,6 +33,7 @@
|
||||
#include <ATen/ops/_scaled_dot_product_flash_attention_backward_native.h>
|
||||
#endif
|
||||
|
||||
+#include <ATen/native/transformers/flash_api_common.h>
|
||||
#ifdef USE_FLASH_ATTENTION
|
||||
// FlashAttention Specific Imports
|
||||
#include <ATen/native/transformers/cuda/flash_attn/flash_api.h>
|
||||
--- /dev/null
|
||||
+++ b/aten/src/ATen/native/transformers/flash_api_common.h
|
||||
@@ -0,0 +1,28 @@
|
||||
+#pragma once
|
||||
+#include <cstdint>
|
||||
+#include <limits>
|
||||
+
|
||||
+#include <ATen/core/Tensor.h>
|
||||
+#include <c10/util/Exception.h>
|
||||
+
|
||||
+#define CHECK_NOSPARSE_CONTIGUOUS_CUDA(TENSOR) \
|
||||
+ TORCH_CHECK(TENSOR.is_cuda(), #TENSOR " must be a CUDA tensor"); \
|
||||
+ TORCH_CHECK(!TENSOR.is_sparse(), #TENSOR " must be a dense tensor"); \
|
||||
+ TORCH_CHECK(TENSOR.is_contiguous());
|
||||
+
|
||||
+#define CHECK_NOSPARSE_LASTCONTIGUOUS_CUDA(TENSOR) \
|
||||
+ TORCH_CHECK(TENSOR.is_cuda(), #TENSOR " must be a CUDA tensor"); \
|
||||
+ TORCH_CHECK(!TENSOR.is_sparse(), #TENSOR " must be a dense tensor"); \
|
||||
+ TORCH_CHECK( \
|
||||
+ TENSOR.stride(-1) == 1, #TENSOR ": last dimension must be contiguous");
|
||||
+
|
||||
+#define CHECK_ALIGNED_PTR(PTR, ALIGNMENT) \
|
||||
+ TORCH_CHECK( \
|
||||
+ uint64_t(PTR) % ALIGNMENT == 0, #PTR " is not correctly aligned")
|
||||
+
|
||||
+#define ASSIGN_CHECK_OVERFLOW(A, B) \
|
||||
+ { \
|
||||
+ A = B; \
|
||||
+ TORCH_CHECK( \
|
||||
+ B < std::numeric_limits<decltype(A)>::max(), #B " overflows"); \
|
||||
+ }
|
||||
--- a/aten/src/ATen/native/transformers/hip/flash_attn/flash_api.h
|
||||
+++ b/aten/src/ATen/native/transformers/hip/flash_attn/flash_api.h
|
||||
@@ -4,28 +4,7 @@
|
||||
#include <ATen/Context.h>
|
||||
#include <ATen/core/Tensor.h>
|
||||
#include <c10/util/Exception.h>
|
||||
-
|
||||
-#define CHECK_NOSPARSE_CONTIGUOUS_CUDA(TENSOR) \
|
||||
- TORCH_CHECK(TENSOR.is_cuda(), #TENSOR " must be a CUDA tensor"); \
|
||||
- TORCH_CHECK(!TENSOR.is_sparse(), #TENSOR " must be a dense tensor"); \
|
||||
- TORCH_CHECK(TENSOR.is_contiguous());
|
||||
-
|
||||
-#define CHECK_NOSPARSE_LASTCONTIGUOUS_CUDA(TENSOR) \
|
||||
- TORCH_CHECK(TENSOR.is_cuda(), #TENSOR " must be a CUDA tensor"); \
|
||||
- TORCH_CHECK(!TENSOR.is_sparse(), #TENSOR " must be a dense tensor"); \
|
||||
- TORCH_CHECK( \
|
||||
- TENSOR.stride(-1) == 1, #TENSOR ": last dimension must be contiguous");
|
||||
-
|
||||
-#define CHECK_ALIGNED_PTR(PTR, ALIGNMENT) \
|
||||
- TORCH_CHECK( \
|
||||
- uint64_t(PTR) % ALIGNMENT == 0, #PTR " is not correctly aligned")
|
||||
-
|
||||
-#define ASSIGN_CHECK_OVERFLOW(A, B) \
|
||||
- { \
|
||||
- A = B; \
|
||||
- TORCH_CHECK( \
|
||||
- B < std::numeric_limits<decltype(A)>::max(), #B " overflows"); \
|
||||
- }
|
||||
+#include <ATen/native/transformers/flash_api_common.h>
|
||||
|
||||
namespace pytorch_flash {
|
||||
|
||||
Reference in New Issue
Block a user