Remove no longer needed const_cast (#3325)

This commit is contained in:
Cheng
2026-03-31 08:10:49 +09:00
committed by GitHub
parent 8a6d28713c
commit 66f58032dc
3 changed files with 12 additions and 19 deletions
+3 -11
View File
@@ -98,9 +98,7 @@ void all_reduce(
size_t block_step;
size_t insize = in.size();
Dtype dt = in.dtype();
// Cub doesn't like const pointers for load (sigh).
void* indata = const_cast<void*>(gpu_ptr<void>(in));
void* indata = gpu_ptr<void>(in);
// Large array so allocate an intermediate and accumulate there
std::tie(blocks, threads, block_step) = get_args(insize, N_READS);
@@ -120,7 +118,7 @@ void all_reduce(
kernel,
blocks,
threads,
static_cast<T*>(indata),
indata,
gpu_ptr<U>(intermediate),
block_step,
insize);
@@ -143,13 +141,7 @@ void all_reduce(
using U = typename cu::ReduceResult<OP, T>::type;
auto kernel = cu::all_reduce<T, U, OP, N_READS>;
encoder.add_kernel_node(
kernel,
blocks,
threads,
static_cast<T*>(indata),
gpu_ptr<U>(out),
block_step,
insize);
kernel, blocks, threads, indata, gpu_ptr<U>(out), block_step, insize);
});
});
}
+2 -6
View File
@@ -282,8 +282,6 @@ void col_reduce_looped(
using OP = MLX_GET_TYPE(reduce_type_tag);
using T = cuda_type_t<MLX_GET_TYPE(type_tag)>;
using U = typename cu::ReduceResult<OP, T>::type;
// Cub doesn't like const pointers for vectorized loads. (sigh)
T* indata = const_cast<T*>(gpu_ptr<T>(in));
constexpr int N_READS = 4;
constexpr int BM = 32;
@@ -296,7 +294,7 @@ void col_reduce_looped(
kernel,
grid,
blocks,
indata,
gpu_ptr<T>(in),
gpu_ptr<U>(out),
static_cast<cu::ColReduceArgs>(args),
out.size() / args.reduction_stride);
@@ -389,8 +387,6 @@ void col_reduce_two_pass(
using OP = MLX_GET_TYPE(reduce_type_tag);
using T = cuda_type_t<MLX_GET_TYPE(type_tag)>;
using U = typename cu::ReduceResult<OP, T>::type;
// Cub doesn't like const pointers for vectorized loads. (sigh)
T* indata = const_cast<T*>(gpu_ptr<T>(in));
constexpr int N_READS = 4;
constexpr int BM = 32;
@@ -403,7 +399,7 @@ void col_reduce_two_pass(
kernel,
grid,
blocks,
indata,
gpu_ptr<T>(in),
gpu_ptr<U>(intermediate),
static_cast<cu::ColReduceArgs>(args),
out.size() / args.reduction_stride);
+7 -2
View File
@@ -268,10 +268,15 @@ void row_reduce_simple(
kernel = cu::row_reduce_simple<T, U, OP, N_READS, 2>;
}
T* indata = const_cast<T*>(gpu_ptr<T>(in));
int size = plan.shape.back();
encoder.add_kernel_node(
kernel, grid, block, indata, gpu_ptr<U>(out), out.size(), size);
kernel,
grid,
block,
gpu_ptr<T>(in),
gpu_ptr<U>(out),
out.size(),
size);
});
});
}