Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
432 changes: 432 additions & 0 deletions benchmarking/optimizer_current_stream.py

Large diffs are not rendered by default.

118 changes: 80 additions & 38 deletions bitsandbytes/backends/cuda/ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -65,6 +65,38 @@ def _setup_ctypes(names, argtypes, restype=None):
[ct.c_void_p] * 4 + [ct.c_int32, ct.c_int32],
)

# 32-bit optimizer update: (g, p, state1, state2, unorm, optimizer scalars, step, lr, gnorm, skip, n, stream)
_setup_ctypes(
[
f"c{name}32bit_grad_{dtype}_with_stream"
for name, dtypes in (
("adam", ("fp32", "fp16", "bf16")),
("momentum", ("32", "16")),
("rmsprop", ("32", "16")),
("lion", ("fp32", "fp16", "bf16")),
("adagrad", ("32", "16")),
("ademamix", ("fp32", "fp16", "bf16")),
)
for dtype in dtypes
],
[ct.c_void_p] * 5 + [ct.c_float] * 8 + [ct.c_int32, ct.c_float, ct.c_float, ct.c_bool, ct.c_int32, ct.c_void_p],
)

# Blockwise 8-bit optimizer update: (p, g, states, scalars, step, lr, maps, absmax, weight decay, gnorm, skip, n, stream)
_setup_ctypes(
[
f"c{name}_8bit_blockwise_grad_{dtype}_with_stream"
for name in ("adam", "momentum", "rmsprop", "lion", "adagrad", "ademamix")
for dtype in ("fp32", "fp16", "bf16")
],
[ct.c_void_p] * 4
+ [ct.c_float] * 5
+ [ct.c_int32, ct.c_float]
+ [ct.c_void_p] * 4
+ [ct.c_float] * 2
+ [ct.c_bool, ct.c_int32, ct.c_void_p],
)


_get_raw_stream = torch._C._cuda_getCurrentRawStream

Expand Down Expand Up @@ -985,73 +1017,73 @@ def _(
"""C FUNCTIONS FOR OPTIMIZERS"""
str2optimizer32bit = {
"adam": (
lib.cadam32bit_grad_fp32,
lib.cadam32bit_grad_fp16,
lib.cadam32bit_grad_bf16,
lib.cadam32bit_grad_fp32_with_stream,
lib.cadam32bit_grad_fp16_with_stream,
lib.cadam32bit_grad_bf16_with_stream,
),
"momentum": (
lib.cmomentum32bit_grad_32,
lib.cmomentum32bit_grad_16,
lib.cmomentum32bit_grad_32_with_stream,
lib.cmomentum32bit_grad_16_with_stream,
),
"rmsprop": (
lib.crmsprop32bit_grad_32,
lib.crmsprop32bit_grad_16,
lib.crmsprop32bit_grad_32_with_stream,
lib.crmsprop32bit_grad_16_with_stream,
),
"lion": (
lib.clion32bit_grad_fp32,
lib.clion32bit_grad_fp16,
lib.clion32bit_grad_bf16,
lib.clion32bit_grad_fp32_with_stream,
lib.clion32bit_grad_fp16_with_stream,
lib.clion32bit_grad_bf16_with_stream,
),
"adagrad": (
lib.cadagrad32bit_grad_32,
lib.cadagrad32bit_grad_16,
lib.cadagrad32bit_grad_32_with_stream,
lib.cadagrad32bit_grad_16_with_stream,
),
"lamb": (
lib.cadam32bit_grad_fp32,
lib.cadam32bit_grad_fp16,
lib.cadam32bit_grad_bf16,
lib.cadam32bit_grad_fp32_with_stream,
lib.cadam32bit_grad_fp16_with_stream,
lib.cadam32bit_grad_bf16_with_stream,
),
"ademamix": (
lib.cademamix32bit_grad_fp32,
lib.cademamix32bit_grad_fp16,
lib.cademamix32bit_grad_bf16,
lib.cademamix32bit_grad_fp32_with_stream,
lib.cademamix32bit_grad_fp16_with_stream,
lib.cademamix32bit_grad_bf16_with_stream,
),
"lars": (
lib.cmomentum32bit_grad_32,
lib.cmomentum32bit_grad_16,
lib.cmomentum32bit_grad_32_with_stream,
lib.cmomentum32bit_grad_16_with_stream,
),
}

str2optimizer8bit_blockwise = {
"adam": (
lib.cadam_8bit_blockwise_grad_fp32,
lib.cadam_8bit_blockwise_grad_fp16,
lib.cadam_8bit_blockwise_grad_bf16,
lib.cadam_8bit_blockwise_grad_fp32_with_stream,
lib.cadam_8bit_blockwise_grad_fp16_with_stream,
lib.cadam_8bit_blockwise_grad_bf16_with_stream,
),
"momentum": (
lib.cmomentum_8bit_blockwise_grad_fp32,
lib.cmomentum_8bit_blockwise_grad_fp16,
lib.cmomentum_8bit_blockwise_grad_bf16,
lib.cmomentum_8bit_blockwise_grad_fp32_with_stream,
lib.cmomentum_8bit_blockwise_grad_fp16_with_stream,
lib.cmomentum_8bit_blockwise_grad_bf16_with_stream,
),
"rmsprop": (
lib.crmsprop_8bit_blockwise_grad_fp32,
lib.crmsprop_8bit_blockwise_grad_fp16,
lib.crmsprop_8bit_blockwise_grad_bf16,
lib.crmsprop_8bit_blockwise_grad_fp32_with_stream,
lib.crmsprop_8bit_blockwise_grad_fp16_with_stream,
lib.crmsprop_8bit_blockwise_grad_bf16_with_stream,
),
"lion": (
lib.clion_8bit_blockwise_grad_fp32,
lib.clion_8bit_blockwise_grad_fp16,
lib.clion_8bit_blockwise_grad_bf16,
lib.clion_8bit_blockwise_grad_fp32_with_stream,
lib.clion_8bit_blockwise_grad_fp16_with_stream,
lib.clion_8bit_blockwise_grad_bf16_with_stream,
),
"adagrad": (
lib.cadagrad_8bit_blockwise_grad_fp32,
lib.cadagrad_8bit_blockwise_grad_fp16,
lib.cadagrad_8bit_blockwise_grad_bf16,
lib.cadagrad_8bit_blockwise_grad_fp32_with_stream,
lib.cadagrad_8bit_blockwise_grad_fp16_with_stream,
lib.cadagrad_8bit_blockwise_grad_bf16_with_stream,
),
"ademamix": (
lib.cademamix_8bit_blockwise_grad_fp32,
lib.cademamix_8bit_blockwise_grad_fp16,
lib.cademamix_8bit_blockwise_grad_bf16,
lib.cademamix_8bit_blockwise_grad_fp32_with_stream,
lib.cademamix_8bit_blockwise_grad_fp16_with_stream,
lib.cademamix_8bit_blockwise_grad_bf16_with_stream,
),
}

Expand Down Expand Up @@ -1092,7 +1124,11 @@ def _optimizer_update_32bit_impl(
f"Gradient+optimizer bit data type combination not supported: grad {g.dtype}, optimizer {state1.dtype}",
)

is_paged = getattr(state1, "is_paged", False) or (state2 is not None and getattr(state2, "is_paged", False))

with _cuda_device_of(g):
# Managed-state prefetches use stream 0, so keep actual paged updates ordered behind them.
stream = None if is_paged else _get_raw_stream(g.device.index)
optim_func(
get_ptr(g),
get_ptr(p),
Expand All @@ -1112,6 +1148,7 @@ def _optimizer_update_32bit_impl(
ct.c_float(gnorm_scale),
ct.c_bool(skip_zeros),
ct.c_int32(g.numel()),
ct.c_void_p(stream),
)


Expand Down Expand Up @@ -1184,7 +1221,11 @@ def _optimizer_update_8bit_blockwise_impl(
f"Unsupported gradient dtype: {g.dtype}. Supported dtypes: torch.float32, torch.float16, torch.bfloat16"
)

is_paged = getattr(state1, "is_paged", False) or (state2 is not None and getattr(state2, "is_paged", False))

with _cuda_device_of(g):
# Managed-state prefetches use stream 0, so keep actual paged updates ordered behind them.
stream = None if is_paged else _get_raw_stream(g.device.index)
optimizer_fn(
get_ptr(p),
get_ptr(g),
Expand All @@ -1205,6 +1246,7 @@ def _optimizer_update_8bit_blockwise_impl(
ct.c_float(gnorm_scale),
ct.c_bool(skip_zeros),
ct.c_int32(g.numel()),
ct.c_void_p(stream),
)


Expand Down
3 changes: 2 additions & 1 deletion bitsandbytes/optim/optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -332,7 +332,8 @@ def step(self, closure=None):

self.prefetch_state(p)
self.update_step(group, p, gindex, pindex)
sync_gpu(p)
if self.is_paged or p.device.type != "cuda" or torch.version.hip is not None:
sync_gpu(p)
if self.is_paged and p is not None:
# all paged operations are asynchronous, we need
# to sync to make sure all tensors are in the right state
Expand Down
2 changes: 2 additions & 0 deletions csrc/compat.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -58,6 +58,7 @@ using bnb_error_t = hipError_t;
#define BNB_DEVICE_MALLOC(p, s) hipMalloc(p, s)
#define BNB_DEVICE_FREE(p) hipFree(p)
#define BNB_DEVICE_MEMSET(p, v, s) hipMemset(p, v, s)
#define BNB_DEVICE_MEMSET_ASYNC(p, v, s, stream) hipMemsetAsync(p, v, s, stream)

#else // CUDA

Expand All @@ -70,6 +71,7 @@ using bnb_error_t = cudaError_t;
#define BNB_DEVICE_MALLOC(p, s) cudaMalloc(p, s)
#define BNB_DEVICE_FREE(p) cudaFree(p)
#define BNB_DEVICE_MEMSET(p, v, s) cudaMemset(p, v, s)
#define BNB_DEVICE_MEMSET_ASYNC(p, v, s, stream) cudaMemsetAsync(p, v, s, stream)

#endif

Expand Down
36 changes: 19 additions & 17 deletions csrc/ops.cu
Original file line number Diff line number Diff line change
Expand Up @@ -97,21 +97,21 @@ template <typename T, int OPTIMIZER>
void optimizer32bit(
T* g, T* p, float* state1, float* state2, float* unorm, float max_unorm, float param_norm, const float beta1,
const float beta2, const float beta3, const float alpha, const float eps, const float weight_decay, const int step,
const float lr, const float gnorm_scale, bool skip_zeros, const int n
const float lr, const float gnorm_scale, bool skip_zeros, const int n, bnb_stream_t stream
) {
int num_blocks = n / 4096;
num_blocks = n % 4096 == 0 ? num_blocks : num_blocks + 1;
switch (OPTIMIZER) {
case ADAM:
case ADEMAMIX:
if (max_unorm > 0.0f) {
BNB_CHECK_RETURN(BNB_DEVICE_MEMSET(unorm, 0, 1 * sizeof(float)));
kPreconditionOptimizer32bit2State<T, OPTIMIZER, 4096, 8><<<num_blocks, 512>>>(
BNB_CHECK_RETURN(BNB_DEVICE_MEMSET_ASYNC(unorm, 0, 1 * sizeof(float), stream));
kPreconditionOptimizer32bit2State<T, OPTIMIZER, 4096, 8><<<num_blocks, 512, 0, stream>>>(
g, p, state1, state2, unorm, beta1, beta2, eps, weight_decay, step, lr, gnorm_scale, n
);
BNB_CHECK_RETURN(BNB_PEEK_LAST_ERROR());
}
kOptimizer32bit2State<T, OPTIMIZER><<<num_blocks, 1024>>>(
kOptimizer32bit2State<T, OPTIMIZER><<<num_blocks, 1024, 0, stream>>>(
g, p, state1, state2, unorm, max_unorm, param_norm, beta1, beta2, beta3, alpha, eps, weight_decay, step, lr,
gnorm_scale, skip_zeros, n
);
Expand All @@ -121,30 +121,32 @@ void optimizer32bit(
case RMSPROP:
case ADAGRAD:
if (max_unorm > 0.0f) {
BNB_CHECK_RETURN(BNB_DEVICE_MEMSET(unorm, 0, 1 * sizeof(float)));
kPreconditionOptimizer32bit1State<T, OPTIMIZER, 4096, 8>
<<<num_blocks, 512>>>(g, p, state1, unorm, beta1, beta2, eps, weight_decay, step, lr, gnorm_scale, n);
BNB_CHECK_RETURN(BNB_DEVICE_MEMSET_ASYNC(unorm, 0, 1 * sizeof(float), stream));
kPreconditionOptimizer32bit1State<T, OPTIMIZER, 4096, 8><<<num_blocks, 512, 0, stream>>>(
g, p, state1, unorm, beta1, beta2, eps, weight_decay, step, lr, gnorm_scale, n
);
BNB_CHECK_RETURN(BNB_PEEK_LAST_ERROR());
}

kOptimizer32bit1State<T, OPTIMIZER><<<num_blocks, 1024>>>(
kOptimizer32bit1State<T, OPTIMIZER><<<num_blocks, 1024, 0, stream>>>(
g, p, state1, unorm, max_unorm, param_norm, beta1, beta2, eps, weight_decay, step, lr, gnorm_scale,
skip_zeros, n
);
BNB_CHECK_RETURN(BNB_PEEK_LAST_ERROR());
break;
case LION:
// in lion, the momentum update after the parameter update
kOptimizer32bit1State<T, OPTIMIZER><<<num_blocks, 1024>>>(
kOptimizer32bit1State<T, OPTIMIZER><<<num_blocks, 1024, 0, stream>>>(
g, p, state1, unorm, max_unorm, param_norm, beta1, beta2, eps, weight_decay, step, lr, gnorm_scale,
skip_zeros, n
);
BNB_CHECK_RETURN(BNB_PEEK_LAST_ERROR());

if (max_unorm > 0.0f) {
BNB_CHECK_RETURN(BNB_DEVICE_MEMSET(unorm, 0, 1 * sizeof(float)));
kPreconditionOptimizer32bit1State<T, OPTIMIZER, 4096, 8>
<<<num_blocks, 512>>>(g, p, state1, unorm, beta1, beta2, eps, weight_decay, step, lr, gnorm_scale, n);
BNB_CHECK_RETURN(BNB_DEVICE_MEMSET_ASYNC(unorm, 0, 1 * sizeof(float), stream));
kPreconditionOptimizer32bit1State<T, OPTIMIZER, 4096, 8><<<num_blocks, 512, 0, stream>>>(
g, p, state1, unorm, beta1, beta2, eps, weight_decay, step, lr, gnorm_scale, n
);
BNB_CHECK_RETURN(BNB_PEEK_LAST_ERROR());
}
break;
Expand All @@ -160,7 +162,7 @@ template <typename T, int OPTIMIZER>
void optimizerStatic8bitBlockwise(
T* p, T* g, unsigned char* state1, unsigned char* state2, float beta1, float beta2, float beta3, float alpha,
float eps, int step, float lr, float* quantiles1, float* quantiles2, float* absmax1, float* absmax2,
float weight_decay, const float gnorm_scale, bool skip_zeros, int n
float weight_decay, const float gnorm_scale, bool skip_zeros, int n, bnb_stream_t stream
) {

int num_blocks = 0;
Expand All @@ -170,7 +172,7 @@ void optimizerStatic8bitBlockwise(
num_blocks = n / BLOCKSIZE_2STATE;
num_blocks = n % BLOCKSIZE_2STATE == 0 ? num_blocks : num_blocks + 1;
kOptimizerStatic8bit2StateBlockwise<T, OPTIMIZER, BLOCKSIZE_2STATE, NUM_2STATE>
<<<num_blocks, BLOCKSIZE_2STATE / NUM_2STATE>>>(
<<<num_blocks, BLOCKSIZE_2STATE / NUM_2STATE, 0, stream>>>(
p, g, state1, state2, beta1, beta2, beta3, alpha, eps, step, lr, quantiles1, quantiles2, absmax1,
absmax2, weight_decay, gnorm_scale, skip_zeros, n
);
Expand All @@ -183,7 +185,7 @@ void optimizerStatic8bitBlockwise(
num_blocks = n / BLOCKSIZE_1STATE;
num_blocks = n % BLOCKSIZE_1STATE == 0 ? num_blocks : num_blocks + 1;
kOptimizerStatic8bit1StateBlockwise<T, OPTIMIZER, BLOCKSIZE_1STATE, NUM_1STATE>
<<<num_blocks, BLOCKSIZE_1STATE / NUM_1STATE>>>(
<<<num_blocks, BLOCKSIZE_1STATE / NUM_1STATE, 0, stream>>>(
p, g, state1, beta1, beta2, eps, step, lr, quantiles1, absmax1, weight_decay, gnorm_scale, skip_zeros, n
);
BNB_CHECK_RETURN(BNB_PEEK_LAST_ERROR());
Expand Down Expand Up @@ -568,7 +570,7 @@ template void dequantizeBlockwise<bnb_bfloat16, NF4>(
gtype * g, gtype * p, float* state1, float* state2, float* unorm, float max_unorm, float param_norm, \
const float beta1, const float beta2, const float beta3, const float alpha, const float eps, \
const float weight_decay, const int step, const float lr, const float gnorm_scale, const bool skip_zeros, \
const int n \
const int n, bnb_stream_t stream \
);

MAKE_optimizer32bit(ADAM, half) MAKE_optimizer32bit(ADAM, float) MAKE_optimizer32bit(ADAM, bnb_bfloat16) MAKE_optimizer32bit(MOMENTUM, half) MAKE_optimizer32bit(MOMENTUM, float) MAKE_optimizer32bit(MOMENTUM, bnb_bfloat16) MAKE_optimizer32bit(RMSPROP, half) MAKE_optimizer32bit(RMSPROP, float) MAKE_optimizer32bit(RMSPROP, bnb_bfloat16) MAKE_optimizer32bit(
Expand All @@ -579,7 +581,7 @@ MAKE_optimizer32bit(ADAM, half) MAKE_optimizer32bit(ADAM, float) MAKE_optimizer3
template void optimizerStatic8bitBlockwise<gtype, optim_name>( \
gtype * p, gtype * g, unsigned char* state1, unsigned char* state2, float beta1, float beta2, float beta3, \
float alpha, float eps, int step, float lr, float* quantiles1, float* quantiles2, float* absmax1, \
float* absmax2, float weight_decay, const float gnorm_scale, bool skip_zeros, int n \
float* absmax2, float weight_decay, const float gnorm_scale, bool skip_zeros, int n, bnb_stream_t stream \
);

MAKE_optimizerStatic8bitBlockwise(half, ADAM);
Expand Down
4 changes: 2 additions & 2 deletions csrc/ops.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -105,14 +105,14 @@ template <typename T, int OPTIMIZER>
void optimizer32bit(
T* g, T* p, float* state1, float* state2, float* unorm, float max_unorm, float param_norm, float beta1, float beta2,
float beta3, float alpha, float eps, float weight_decay, int step, float lr, const float gnorm_scale,
bool skip_zeros, int n
bool skip_zeros, int n, bnb_stream_t stream
);

template <typename T, int OPTIMIZER>
void optimizerStatic8bitBlockwise(
T* p, T* g, unsigned char* state1, unsigned char* state2, float beta1, float beta2, float beta3, float alpha,
float eps, int step, float lr, float* quantiles1, float* quantiles2, float* absmax1, float* absmax2,
float weight_decay, const float gnorm_scale, bool skip_zeros, int n
float weight_decay, const float gnorm_scale, bool skip_zeros, int n, bnb_stream_t stream
);

void gemmex(
Expand Down
Loading