Fix torch compile issues (#3028)
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update pyproject.toml
* Delete .gitignore
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update _utils.py
* Update pyproject.toml
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update chat_templates.py
* Seasame force float16 / float32
* Fix Seasame
* Update loader.py
* Update vision.py
* Update vision.py
* Update vision.py
* Update loader.py
* is_multimodal
* Update loader.py
* Update loader.py
* Update loader.py
* Update loader.py
* Update vision.py
* Update vision.py
* Update vision.py
* UNSLOTH_DISABLE_STATIC_GENERATION
* Update vision.py
* Auto vision detection
* Sesame
* Whisper
* Update loader.py
* Update loader.py
* Update loader.py
* Update mapper.py
* Update vision.py
* Update vision.py
* Update vision.py
* Update vision.py
* Update vision.py
* Update vision.py
* Update loader.py
* Update loader.py
* Update loader.py
* Update loader.py
* Update _utils.py
* Update rl.py
* versioning
* Update rl.py
* Update rl.py
* Update rl.py
* Update rl.py
* Update rl.py
* logging
* Update pyproject.toml
* Update rl.py
* versioning
* Update rl.py
* Update rl.py
* Update rl_replacements.py
* Update rl_replacements.py
* Update rl.py
* Update rl_replacements.py
* Update rl_replacements.py
* logits / temperature
* Update rl_replacements.py
* Update pyproject.toml
* Update rl_replacements.py
* Update rl_replacements.py
* Debugging only
* Update llama.py
* Update llama.py
* Update rl_replacements.py
* Update rl_replacements.py
* Update rl_replacements.py
* Update rl_replacements.py
* Update rl_replacements.py
* Generic efficient GRPO
* Update rl_replacements.py
* Update rl_replacements.py
* Remove debugging
* Update rl_replacements.py
* Update rl_replacements.py
* Update vision.py
* Update llama.py
* Update rl_replacements.py
* versioning
* Update _utils.py
* Update vision.py
* Update mapper.py
* Update loader.py
* Update mapper.py
* Update vision.py
* Update loader.py
* Update vision.py
* Update loader.py
* Update _utils.py
* Update vision.py
* gradient checkpointing
* Gemma 3N fixes
* Update loader.py
* Versioning
* Gemma 3N fixes
* Update vision.py
* Update vision.py
* Update loader.py
* Update vision.py
* Fix setup.py
* setup.py
* Prints
* Update setup.py
* Update setup.py
* Update setup.py
* Update pyproject.toml
* Update pyproject.toml
* Update pyproject.toml
* Update pyproject.toml
* Update pyproject.toml
* Update pyproject.toml
* Update vision.py
* Update vision.py
* Update pyproject.toml
* Update vision.py
* Update _utils.py
* Update __init__.py
* Update __init__.py
* Small fixes
* Update vision.py
* Update vision.py
* versioning
* Update __init__.py
* Update llama.py
* Update rl.py
* Update rl.py
* Update _utils.py
* Update vision.py
* Update vision.py
* compiler stance
* Update _utils.py
* Update pyproject.toml
* Update pyproject.toml
* Update rl_replacements.py
* Update rl_replacements.py
* Update rl_replacements.py
* Update rl_replacements.py
* Update rl.py
* Update rl_replacements.py
* Update rl_replacements.py
* Update rl_replacements.py
* Update rl_replacements.py
* Update rl_replacements.py
* Update rl_replacements.py
* Update rl_replacements.py
* Revert "Revert "Add Qwen2.5-VL-32B-Instruct mapping to fix quantized model me…" (#2990)
This reverts commit 4021da634a.
* skip_guard_eval_unsafe fix
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update synthetic.py
* Update llama.py
* Update llama.py
* Fix `quantization_method`
* versioning
* Update _utils.py
* Update _utils.py
* Update _utils.py
* check stride
* Cleanup
* Update rope_embedding.py
* Update gemma2.py
* Fix `set_stance`
This commit is contained in:
parent
f2ef5bd16b
commit
d27e4e44d1
8 changed files with 52 additions and 52 deletions
|
|
@ -37,7 +37,7 @@ triton = [
|
|||
]
|
||||
|
||||
huggingface = [
|
||||
"unsloth_zoo>=2025.7.8",
|
||||
"unsloth_zoo>=2025.7.9",
|
||||
"packaging",
|
||||
"tyro",
|
||||
"transformers>=4.51.3,!=4.47.0,!=4.52.0,!=4.52.1,!=4.52.2,!=4.52.3,!=4.53.0",
|
||||
|
|
@ -381,7 +381,7 @@ colab-ampere-torch220 = [
|
|||
"flash-attn>=2.6.3",
|
||||
]
|
||||
colab-new = [
|
||||
"unsloth_zoo>=2025.7.8",
|
||||
"unsloth_zoo>=2025.7.9",
|
||||
"packaging",
|
||||
"tyro",
|
||||
"transformers>=4.51.3,!=4.47.0,!=4.52.0,!=4.52.1,!=4.52.2,!=4.52.3,!=4.53.0",
|
||||
|
|
|
|||
|
|
@ -107,7 +107,7 @@ _cross_entropy_forward = triton.heuristics(
|
|||
|
||||
def _chunked_cross_entropy_forward(
|
||||
logits_ptr ,
|
||||
logits_row_stride ,
|
||||
logits_row_stride : tl.constexpr,
|
||||
loss_ptr ,
|
||||
logsumexp_ptr ,
|
||||
labels_ptr ,
|
||||
|
|
@ -191,9 +191,9 @@ _chunked_cross_entropy_forward = triton.heuristics(
|
|||
|
||||
def _cross_entropy_backward(
|
||||
logits_ptr ,
|
||||
logits_row_stride ,
|
||||
logits_row_stride : tl.constexpr,
|
||||
dloss_ptr ,
|
||||
dloss_row_stride ,
|
||||
dloss_row_stride : tl.constexpr,
|
||||
logsumexp_ptr ,
|
||||
labels_ptr ,
|
||||
VOCAB_SIZE : tl.constexpr,
|
||||
|
|
|
|||
|
|
@ -19,9 +19,9 @@ from .utils import calculate_settings, torch_gpu_device
|
|||
|
||||
@triton.jit
|
||||
def _rms_layernorm_forward(
|
||||
Y, Y_row_stride,
|
||||
X, X_row_stride,
|
||||
W, W_row_stride,
|
||||
Y, Y_row_stride : tl.constexpr,
|
||||
X, X_row_stride : tl.constexpr,
|
||||
W, W_row_stride : tl.constexpr,
|
||||
r, r_row_stride : tl.constexpr,
|
||||
n_cols : tl.constexpr,
|
||||
eps : tl.constexpr,
|
||||
|
|
@ -54,10 +54,10 @@ pass
|
|||
|
||||
|
||||
def _rms_layernorm_backward(
|
||||
dY, dY_row_stride,
|
||||
dX, dX_row_stride,
|
||||
X, X_row_stride,
|
||||
W, W_row_stride,
|
||||
dY, dY_row_stride : tl.constexpr,
|
||||
dX, dX_row_stride : tl.constexpr,
|
||||
X, X_row_stride : tl.constexpr,
|
||||
W, W_row_stride : tl.constexpr,
|
||||
r, r_row_stride : tl.constexpr,
|
||||
# dW, dW_row_stride,
|
||||
n_cols : tl.constexpr,
|
||||
|
|
@ -106,9 +106,9 @@ _rms_layernorm_backward = triton.heuristics(
|
|||
|
||||
@triton.jit
|
||||
def _gemma_rms_layernorm_forward(
|
||||
Y, Y_row_stride,
|
||||
X, X_row_stride,
|
||||
W, W_row_stride,
|
||||
Y, Y_row_stride : tl.constexpr,
|
||||
X, X_row_stride : tl.constexpr,
|
||||
W, W_row_stride : tl.constexpr,
|
||||
r, r_row_stride : tl.constexpr,
|
||||
n_cols : tl.constexpr,
|
||||
eps : tl.constexpr,
|
||||
|
|
|
|||
|
|
@ -19,9 +19,9 @@ from .utils import calculate_settings, torch_gpu_device
|
|||
ROPE_GROUP_SIZE : int = 4
|
||||
|
||||
def _rope_embedding(
|
||||
Q, Q_row_stride,
|
||||
cos, cos_row_stride,
|
||||
sin, sin_row_stride,
|
||||
Q, Q_row_stride: tl.constexpr,
|
||||
cos, cos_row_stride: tl.constexpr,
|
||||
sin, sin_row_stride: tl.constexpr,
|
||||
seqlen,
|
||||
head_dim : tl.constexpr,
|
||||
n_heads : tl.constexpr,
|
||||
|
|
|
|||
|
|
@ -124,7 +124,7 @@ if DEVICE_TYPE == "xpu":
|
|||
(index := torch.xpu.device(i).idx) : ctypes.c_void_p(torch._C._xpu_getCurrentRawStream(index))
|
||||
for i in range(DEVICE_COUNT)
|
||||
}
|
||||
XPU_STREAMS = [None] * (max(_XPU_STREAMS.keys()) + 1)
|
||||
XPU_STREAMS = [None] * (max(_XPU_STREAMS.keys()) + 1)
|
||||
WEIGHT_BUFFERS = [None] * (max(_XPU_STREAMS.keys()) + 1)
|
||||
ABSMAX_BUFFERS = [None] * (max(_XPU_STREAMS.keys()) + 1)
|
||||
for k, v in _XPU_STREAMS.items():
|
||||
|
|
@ -143,7 +143,7 @@ else:
|
|||
for k, v in _CUDA_STREAMS.items(): CUDA_STREAMS[k] = v
|
||||
CUDA_STREAMS = tuple(CUDA_STREAMS)
|
||||
del _CUDA_STREAMS
|
||||
|
||||
pass
|
||||
|
||||
# Bitsandbytes operations
|
||||
ctypes_c_int = ctypes.c_int
|
||||
|
|
@ -172,12 +172,16 @@ else:
|
|||
cdequantize_blockwise_bf16_nf4 = bnb.functional.lib.cdequantize_blockwise_bf16_nf4
|
||||
cgemm_4bit_inference_naive_fp16 = bnb.functional.lib.cgemm_4bit_inference_naive_fp16
|
||||
cgemm_4bit_inference_naive_bf16 = bnb.functional.lib.cgemm_4bit_inference_naive_bf16
|
||||
pass
|
||||
|
||||
torch_mm = torch.mm
|
||||
torch_mv = torch.mv
|
||||
torch_matmul = torch.matmul
|
||||
torch_addmm = torch.addmm
|
||||
torch_empty = torch.empty
|
||||
torch_float32 = torch.float32
|
||||
torch_float16 = torch.float16
|
||||
torch_bfloat16 = torch.bfloat16
|
||||
|
||||
def QUANT_STATE(W): return getattr(W, "quant_state", None)
|
||||
|
||||
|
|
@ -283,7 +287,7 @@ if DEVICE_TYPE == "xpu" and HAS_XPU_STREAM:
|
|||
else:
|
||||
assert(out.shape == shape)
|
||||
assert(out.dtype == dtype)
|
||||
out_absmax = torch_empty(n_elements_absmax, dtype = torch.float32, device = device, requires_grad = False)
|
||||
out_absmax = torch_empty(n_elements_absmax, dtype = torch_float32, device = device, requires_grad = False)
|
||||
pass
|
||||
|
||||
# NF4 dequantization of statistics
|
||||
|
|
@ -296,7 +300,7 @@ if DEVICE_TYPE == "xpu" and HAS_XPU_STREAM:
|
|||
out_absmax += offset
|
||||
|
||||
# Dequantize W
|
||||
fx = cdequantize_blockwise_fp16_nf4 if dtype == torch.float16 else \
|
||||
fx = cdequantize_blockwise_fp16_nf4 if dtype == torch_float16 else \
|
||||
cdequantize_blockwise_bf16_nf4
|
||||
fx(get_ptr(None), get_ptr(W), ptr_out_absmax, get_ptr(out),
|
||||
ctypes_c_int(blocksize), ctypes_c_int(out.numel()), XPU_STREAM,)
|
||||
|
|
@ -346,7 +350,7 @@ elif DEVICE_TYPE == "cuda" and HAS_CUDA_STREAM:
|
|||
ABSMAX_BUFFER = ABSMAX_BUFFERS[device_index]
|
||||
if WEIGHT_BUFFER is None:
|
||||
WEIGHT_BUFFERS[device_index] = WEIGHT_BUFFER = torch_empty(size, dtype = dtype, device = device, requires_grad = False)
|
||||
ABSMAX_BUFFERS[device_index] = ABSMAX_BUFFER = torch_empty(n_elements_absmax, dtype = torch.float32, device = device, requires_grad = False)
|
||||
ABSMAX_BUFFERS[device_index] = ABSMAX_BUFFER = torch_empty(n_elements_absmax, dtype = torch_float32, device = device, requires_grad = False)
|
||||
|
||||
if size > WEIGHT_BUFFER.numel(): WEIGHT_BUFFER.resize_(size)
|
||||
if n_elements_absmax > ABSMAX_BUFFER.numel(): ABSMAX_BUFFER.resize_(n_elements_absmax)
|
||||
|
|
@ -359,7 +363,7 @@ elif DEVICE_TYPE == "cuda" and HAS_CUDA_STREAM:
|
|||
else:
|
||||
assert(out.shape == shape)
|
||||
assert(out.dtype == dtype)
|
||||
out_absmax = torch_empty(n_elements_absmax, dtype = torch.float32, device = device, requires_grad = False)
|
||||
out_absmax = torch_empty(n_elements_absmax, dtype = torch_float32, device = device, requires_grad = False)
|
||||
pass
|
||||
|
||||
# NF4 dequantization of statistics
|
||||
|
|
@ -372,7 +376,7 @@ elif DEVICE_TYPE == "cuda" and HAS_CUDA_STREAM:
|
|||
out_absmax += offset
|
||||
|
||||
# Dequantize W
|
||||
fx = cdequantize_blockwise_fp16_nf4 if dtype == torch.float16 else \
|
||||
fx = cdequantize_blockwise_fp16_nf4 if dtype == torch_float16 else \
|
||||
cdequantize_blockwise_bf16_nf4
|
||||
fx(get_ptr(None), get_ptr(W), ptr_out_absmax, get_ptr(out),
|
||||
ctypes_c_int(blocksize), ctypes_c_int(out.numel()), CUDA_STREAM,)
|
||||
|
|
@ -413,7 +417,7 @@ else:
|
|||
else:
|
||||
assert(out.shape == shape)
|
||||
assert(out.dtype == dtype)
|
||||
out_absmax = torch_empty(n_elements_absmax, dtype = torch.float32, device = device, requires_grad = False)
|
||||
out_absmax = torch_empty(n_elements_absmax, dtype = torch_float32, device = device, requires_grad = False)
|
||||
|
||||
# Do dequantization
|
||||
ptr_out_absmax = get_ptr(out_absmax)
|
||||
|
|
@ -423,7 +427,7 @@ else:
|
|||
)
|
||||
out_absmax += offset
|
||||
|
||||
fx = cdequantize_blockwise_fp16_nf4 if dtype == torch.float16 else \
|
||||
fx = cdequantize_blockwise_fp16_nf4 if dtype == torch_float16 else \
|
||||
cdequantize_blockwise_bf16_nf4
|
||||
fx(get_ptr(None), get_ptr(W), ptr_out_absmax, get_ptr(out),
|
||||
ctypes_c_int(blocksize), ctypes_c_int(out.numel()),)
|
||||
|
|
@ -488,7 +492,7 @@ if DEVICE_TYPE == "xpu" and HAS_XPU_STREAM:
|
|||
ldb = ctypes_c_int32(ldb)
|
||||
ldc = ctypes_c_int32(ldc)
|
||||
|
||||
df = torch_empty(absmax.shape, dtype = torch.float32, device = device)
|
||||
df = torch_empty(absmax.shape, dtype = torch_float32, device = device)
|
||||
with torch_gpu_device(device):
|
||||
cdequantize_blockwise_fp32(
|
||||
get_ptr(code2), get_ptr(absmax), get_ptr(absmax2), get_ptr(df),
|
||||
|
|
@ -497,7 +501,7 @@ if DEVICE_TYPE == "xpu" and HAS_XPU_STREAM:
|
|||
df += offset
|
||||
absmax = df
|
||||
|
||||
fx = cgemm_4bit_inference_naive_fp16 if dtype == torch.float16 else \
|
||||
fx = cgemm_4bit_inference_naive_fp16 if dtype == torch_float16 else \
|
||||
cgemm_4bit_inference_naive_bf16
|
||||
|
||||
blocksize = ctypes_c_int32(blocksize)
|
||||
|
|
@ -559,7 +563,7 @@ elif DEVICE_TYPE == "cuda" and HAS_CUDA_STREAM:
|
|||
ldb = ctypes_c_int32(ldb)
|
||||
ldc = ctypes_c_int32(ldc)
|
||||
|
||||
df = torch_empty(absmax.shape, dtype = torch.float32, device = device)
|
||||
df = torch_empty(absmax.shape, dtype = torch_float32, device = device)
|
||||
with torch_gpu_device(device):
|
||||
cdequantize_blockwise_fp32(
|
||||
get_ptr(code2), get_ptr(absmax), get_ptr(absmax2), get_ptr(df),
|
||||
|
|
@ -568,8 +572,8 @@ elif DEVICE_TYPE == "cuda" and HAS_CUDA_STREAM:
|
|||
df += offset
|
||||
absmax = df
|
||||
|
||||
fx = cgemm_4bit_inference_naive_fp16 if dtype == torch.float16 else \
|
||||
cgemm_4bit_inference_naive_bf16
|
||||
fx = cgemm_4bit_inference_naive_fp16 if dtype == torch_float16 else \
|
||||
cgemm_4bit_inference_naive_bf16
|
||||
|
||||
blocksize = ctypes_c_int32(blocksize)
|
||||
fx(m, n, k, get_ptr(X), get_ptr(W), get_ptr(absmax), get_ptr(stats), get_ptr(out),
|
||||
|
|
@ -580,7 +584,7 @@ elif DEVICE_TYPE == "cuda" and HAS_CUDA_STREAM:
|
|||
pass
|
||||
else:
|
||||
def fast_gemv(X, W, quant_state, out = None):
|
||||
if quant_state is None: return torch.matmul(X, W, out = out)
|
||||
if quant_state is None: return torch_matmul(X, W, out = out)
|
||||
# For fast X @ W where seq_len == 1
|
||||
# From https://github.com/TimDettmers/bitsandbytes/blob/main/bitsandbytes/functional.py#L1469
|
||||
_, q_len, hd = X.shape
|
||||
|
|
@ -626,7 +630,7 @@ else:
|
|||
ldb = ctypes_c_int32(ldb)
|
||||
ldc = ctypes_c_int32(ldc)
|
||||
|
||||
df = torch_empty(absmax.shape, dtype = torch.float32, device = device)
|
||||
df = torch_empty(absmax.shape, dtype = torch_float32, device = device)
|
||||
cdequantize_blockwise_fp32(
|
||||
get_ptr(code2), get_ptr(absmax), get_ptr(absmax2), get_ptr(df),
|
||||
ctypes_c_int(blocksize2), ctypes_c_int(df.numel()),
|
||||
|
|
@ -634,8 +638,8 @@ else:
|
|||
df += offset
|
||||
absmax = df
|
||||
|
||||
fx = cgemm_4bit_inference_naive_fp16 if dtype == torch.float16 else \
|
||||
cgemm_4bit_inference_naive_bf16
|
||||
fx = cgemm_4bit_inference_naive_fp16 if dtype == torch_float16 else \
|
||||
cgemm_4bit_inference_naive_bf16
|
||||
|
||||
blocksize = ctypes_c_int32(blocksize)
|
||||
fx(m, n, k, get_ptr(X), get_ptr(W), get_ptr(absmax), get_ptr(stats), get_ptr(out),
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@
|
|||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
__version__ = "2025.7.7"
|
||||
__version__ = "2025.7.8"
|
||||
|
||||
__all__ = [
|
||||
"SUPPORTS_BFLOAT16",
|
||||
|
|
@ -431,7 +431,7 @@ if DEVICE_TYPE == "cuda":
|
|||
"Unsloth: If you want to finetune Gemma 2, upgrade flash-attn to version 2.6.3 or higher!\n"\
|
||||
"Newer versions support faster and less memory usage kernels for Gemma 2's attention softcapping!\n"\
|
||||
"To update flash-attn, do the below:\n"\
|
||||
'\npip install --no-deps --upgrade "flash-attn>=2.6.3"'
|
||||
'\npip install --no-deps --no-build-isolation --upgrade "flash-attn>=2.6.3"'
|
||||
)
|
||||
except:
|
||||
print(
|
||||
|
|
|
|||
|
|
@ -59,18 +59,6 @@ pass
|
|||
if HAS_FLASH_ATTENTION_SOFTCAPPING:
|
||||
from flash_attn import flash_attn_func
|
||||
|
||||
# [TODO] We must randomnly use torch.compile?
|
||||
# Gemma 2 uses double RMS Layernorms, so the backward passes should not overwrite the gradients!
|
||||
@torch.compile(fullgraph = False, dynamic = True, options = torch_compile_options)
|
||||
def fast_rms_layernorm_gemma2_compiled(layernorm, X, gemma = True):
|
||||
old_dtype = X.dtype
|
||||
X = X.float()
|
||||
X = X * torch.rsqrt(X.square().mean(-1, keepdim = True) + layernorm.eps) * \
|
||||
(1.0 + layernorm.weight.float())
|
||||
return X.to(old_dtype)
|
||||
pass
|
||||
|
||||
|
||||
# Logit softcapping
|
||||
def Gemma2Attention_fast_forward(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -85,6 +85,12 @@ from unsloth_zoo.vllm_utils import (
|
|||
return_lora_modules,
|
||||
)
|
||||
|
||||
try:
|
||||
torch_compiler_set_stance = torch.compiler.set_stance
|
||||
except:
|
||||
torch_compiler_set_stance = None
|
||||
pass
|
||||
|
||||
def unsloth_base_fast_generate(
|
||||
self,
|
||||
*args,
|
||||
|
|
@ -756,7 +762,8 @@ class FastBaseModel:
|
|||
# Must enable returning logits
|
||||
os.environ["UNSLOTH_RETURN_LOGITS"] = "1"
|
||||
# Turn off skip guards and set stance to default
|
||||
torch.compiler.set_stance(stance = "default", skip_guard_eval_unsafe = False)
|
||||
if torch_compiler_set_stance is not None:
|
||||
torch_compiler_set_stance(stance = "default", skip_guard_eval_unsafe = False)
|
||||
return model
|
||||
pass
|
||||
|
||||
|
|
@ -804,7 +811,8 @@ class FastBaseModel:
|
|||
# Can re-enable not returning logits
|
||||
os.environ["UNSLOTH_RETURN_LOGITS"] = "0"
|
||||
# Turn off skip guards and set stance to default
|
||||
torch.compiler.set_stance(stance = "default", skip_guard_eval_unsafe = False)
|
||||
if torch_compiler_set_stance is not None:
|
||||
torch_compiler_set_stance(stance = "default", skip_guard_eval_unsafe = False)
|
||||
return model
|
||||
pass
|
||||
pass
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue