fix: quantize amax from raw values, not scaled fp8 values
- amax for delayed scale was the quantized max (always ~448), so scale collapsed to 1 - this made fp8 gradients diverge (cosine 0.05) and training stall - stop w/x transpose-quantize amax from polluting the grad scale
This commit is contained in:
+59
-9
@@ -52,8 +52,15 @@ class FP8TensorMeta:
|
|||||||
"idx",
|
"idx",
|
||||||
"x_scale",
|
"x_scale",
|
||||||
"x_scale_inv",
|
"x_scale_inv",
|
||||||
|
"x_history",
|
||||||
|
"x_idx",
|
||||||
"g_scale",
|
"g_scale",
|
||||||
"g_scale_inv",
|
"g_scale_inv",
|
||||||
|
"g_history",
|
||||||
|
"g_idx",
|
||||||
|
"w_init",
|
||||||
|
"x_init",
|
||||||
|
"g_init",
|
||||||
)
|
)
|
||||||
|
|
||||||
def __init__(self, device: torch.device, update_interval: int):
|
def __init__(self, device: torch.device, update_interval: int):
|
||||||
@@ -65,8 +72,43 @@ class FP8TensorMeta:
|
|||||||
self.idx = 0
|
self.idx = 0
|
||||||
self.x_scale = torch.ones(1, device=device, dtype=torch.float32)
|
self.x_scale = torch.ones(1, device=device, dtype=torch.float32)
|
||||||
self.x_scale_inv = torch.ones(1, device=device, dtype=torch.float32)
|
self.x_scale_inv = torch.ones(1, device=device, dtype=torch.float32)
|
||||||
|
self.x_history = torch.ones(update_interval, device=device, dtype=torch.float32)
|
||||||
|
self.x_idx = 0
|
||||||
self.g_scale = torch.ones(1, device=device, dtype=torch.float32)
|
self.g_scale = torch.ones(1, device=device, dtype=torch.float32)
|
||||||
self.g_scale_inv = torch.ones(1, device=device, dtype=torch.float32)
|
self.g_scale_inv = torch.ones(1, device=device, dtype=torch.float32)
|
||||||
|
self.g_history = torch.ones(update_interval, device=device, dtype=torch.float32)
|
||||||
|
self.g_idx = 0
|
||||||
|
self.w_init = False
|
||||||
|
self.x_init = False
|
||||||
|
self.g_init = False
|
||||||
|
|
||||||
|
def init_scale(self, t: torch.Tensor) -> None:
|
||||||
|
"""Immediate scale from the current amax; used on the first call.
|
||||||
|
|
||||||
|
A scale of 1 would underflow small activations/gradients (e4m3 min
|
||||||
|
normal is 2^-6); initialize from the actual amax once, then delayed
|
||||||
|
updates take over.
|
||||||
|
"""
|
||||||
|
amax = t.abs().amax().to(torch.float32).clamp_min(1e-12)
|
||||||
|
self.scale.copy_(amax / E4M3_MAX)
|
||||||
|
self.scale_inv.copy_(E4M3_MAX / amax)
|
||||||
|
self.record(amax)
|
||||||
|
|
||||||
|
def push_x_scale(self, amax: torch.Tensor) -> None:
|
||||||
|
"""Window update for the activation scale (delayed, TE style)."""
|
||||||
|
self.x_history[self.x_idx] = amax.reshape(())
|
||||||
|
self.x_idx = (self.x_idx + 1) % self.x_history.numel()
|
||||||
|
m = self.x_history.max()
|
||||||
|
self.x_scale.copy_(m / E4M3_MAX)
|
||||||
|
self.x_scale_inv.copy_(E4M3_MAX / m)
|
||||||
|
|
||||||
|
def push_g_scale(self, amax: torch.Tensor) -> None:
|
||||||
|
"""Window update for the gradient scale (delayed, TE style)."""
|
||||||
|
self.g_history[self.g_idx] = amax.reshape(())
|
||||||
|
self.g_idx = (self.g_idx + 1) % self.g_history.numel()
|
||||||
|
m = self.g_history.max()
|
||||||
|
self.g_scale.copy_(m / E4M3_MAX)
|
||||||
|
self.g_scale_inv.copy_(E4M3_MAX / m)
|
||||||
|
|
||||||
def record(self, amax: torch.Tensor) -> None:
|
def record(self, amax: torch.Tensor) -> None:
|
||||||
"""Push the latest amax into the ring buffer (device-side copy, no sync)."""
|
"""Push the latest amax into the ring buffer (device-side copy, no sync)."""
|
||||||
@@ -155,13 +197,6 @@ def fp8_autocast(enabled: bool = True, update_interval: int = 16):
|
|||||||
state.update_interval = prev_interval
|
state.update_interval = prev_interval
|
||||||
|
|
||||||
|
|
||||||
def _update_delayed_scale(scale, scale_inv, amax) -> None:
|
|
||||||
"""scale = amax / 448 for the *next* call (device-side, no sync)."""
|
|
||||||
amax_f = amax.reshape(()).to(torch.float32).clamp_min(1e-12)
|
|
||||||
scale.copy_(amax_f / E4M3_MAX)
|
|
||||||
scale_inv.copy_(E4M3_MAX / amax_f)
|
|
||||||
|
|
||||||
|
|
||||||
def fp8_linear_forward(x: torch.Tensor, w: torch.Tensor, bias=None):
|
def fp8_linear_forward(x: torch.Tensor, w: torch.Tensor, bias=None):
|
||||||
"""TE-style scaled fp8 linear forward (called from the aten::linear impl).
|
"""TE-style scaled fp8 linear forward (called from the aten::linear impl).
|
||||||
|
|
||||||
@@ -173,6 +208,15 @@ def fp8_linear_forward(x: torch.Tensor, w: torch.Tensor, bias=None):
|
|||||||
bias = torch.empty(0, device=x.device, dtype=x.dtype)
|
bias = torch.empty(0, device=x.device, dtype=x.dtype)
|
||||||
state = fp8_state()
|
state = fp8_state()
|
||||||
meta = state.get_weight_meta(w)
|
meta = state.get_weight_meta(w)
|
||||||
|
if not meta.w_init:
|
||||||
|
meta.init_scale(w)
|
||||||
|
meta.w_init = True
|
||||||
|
if not meta.x_init:
|
||||||
|
amax = x.abs().amax().to(torch.float32).clamp_min(1e-12)
|
||||||
|
meta.x_history.fill_(amax)
|
||||||
|
meta.x_scale.copy_(amax / E4M3_MAX)
|
||||||
|
meta.x_scale_inv.copy_(E4M3_MAX / amax)
|
||||||
|
meta.x_init = True
|
||||||
amax_x = torch.empty(1, device=x.device, dtype=torch.float32)
|
amax_x = torch.empty(1, device=x.device, dtype=torch.float32)
|
||||||
amax_w = torch.empty(1, device=x.device, dtype=torch.float32)
|
amax_w = torch.empty(1, device=x.device, dtype=torch.float32)
|
||||||
out = linear_forward_scaled(
|
out = linear_forward_scaled(
|
||||||
@@ -187,7 +231,7 @@ def fp8_linear_forward(x: torch.Tensor, w: torch.Tensor, bias=None):
|
|||||||
amax_w,
|
amax_w,
|
||||||
)
|
)
|
||||||
meta.record(amax_w)
|
meta.record(amax_w)
|
||||||
_update_delayed_scale(meta.x_scale, meta.x_scale_inv, amax_x)
|
meta.push_x_scale(amax_x)
|
||||||
return out
|
return out
|
||||||
|
|
||||||
|
|
||||||
@@ -195,6 +239,12 @@ def fp8_linear_backward(g, x, w, masks):
|
|||||||
"""TE-style scaled fp8 linear backward (called from aten::linear_backward)."""
|
"""TE-style scaled fp8 linear backward (called from aten::linear_backward)."""
|
||||||
state = fp8_state()
|
state = fp8_state()
|
||||||
meta = state.get_weight_meta(w)
|
meta = state.get_weight_meta(w)
|
||||||
|
if not meta.g_init:
|
||||||
|
amax = g.abs().amax().to(torch.float32).clamp_min(1e-12)
|
||||||
|
meta.g_history.fill_(amax)
|
||||||
|
meta.g_scale.copy_(amax / E4M3_MAX)
|
||||||
|
meta.g_scale_inv.copy_(E4M3_MAX / amax)
|
||||||
|
meta.g_init = True
|
||||||
amax_g = torch.empty(1, device=g.device, dtype=torch.float32)
|
amax_g = torch.empty(1, device=g.device, dtype=torch.float32)
|
||||||
out = linear_backward_scaled(
|
out = linear_backward_scaled(
|
||||||
g,
|
g,
|
||||||
@@ -209,7 +259,7 @@ def fp8_linear_backward(g, x, w, masks):
|
|||||||
meta.x_scale_inv,
|
meta.x_scale_inv,
|
||||||
amax_g,
|
amax_g,
|
||||||
)
|
)
|
||||||
_update_delayed_scale(meta.g_scale, meta.g_scale_inv, amax_g)
|
meta.push_g_scale(amax_g)
|
||||||
return out
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
+15
-12
@@ -154,9 +154,9 @@ __global__ void quantize_kernel(const __nv_bfloat16* __restrict__ src,
|
|||||||
int64_t i = blockIdx.x * (int64_t)blockDim.x + threadIdx.x;
|
int64_t i = blockIdx.x * (int64_t)blockDim.x + threadIdx.x;
|
||||||
float amax = 0.f;
|
float amax = 0.f;
|
||||||
if (i < n) {
|
if (i < n) {
|
||||||
float v = __bfloat162float(src[i]) * *scale_inv;
|
float raw = __bfloat162float(src[i]);
|
||||||
dst[i] = cast_fp8<T8>(v);
|
dst[i] = cast_fp8<T8>(raw * *scale_inv);
|
||||||
amax = fabsf(v);
|
amax = fabsf(raw);
|
||||||
}
|
}
|
||||||
for (int off = 16; off; off >>= 1)
|
for (int off = 16; off; off >>= 1)
|
||||||
amax = fmaxf(amax, __shfl_xor_sync(0xffffffffu, amax, off));
|
amax = fmaxf(amax, __shfl_xor_sync(0xffffffffu, amax, off));
|
||||||
@@ -182,9 +182,9 @@ __global__ void transpose_quantize_kernel(
|
|||||||
float amax = 0.f;
|
float amax = 0.f;
|
||||||
for (int j = 0; j < 32; j += 8) {
|
for (int j = 0; j < 32; j += 8) {
|
||||||
if (x < cols && y + j < rows) {
|
if (x < cols && y + j < rows) {
|
||||||
float v = __bfloat162float(src[(y + j) * cols + x]) * *scale_inv;
|
float raw = __bfloat162float(src[(y + j) * cols + x]);
|
||||||
tile[threadIdx.y + j][threadIdx.x] = cast_fp8<T8>(v);
|
tile[threadIdx.y + j][threadIdx.x] = cast_fp8<T8>(raw * *scale_inv);
|
||||||
amax = fmaxf(amax, fabsf(v));
|
amax = fmaxf(amax, fabsf(raw));
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
__syncthreads();
|
__syncthreads();
|
||||||
@@ -382,6 +382,9 @@ std::tuple<torch::Tensor, torch::Tensor, torch::Tensor> fp8_linear_backward_scal
|
|||||||
auto gt8 = masks[1] ? torch::empty({n, m}, fp8_options) : torch::Tensor();
|
auto gt8 = masks[1] ? torch::empty({n, m}, fp8_options) : torch::Tensor();
|
||||||
auto wt8 = masks[0] ? torch::empty({k, n}, fp8_options) : torch::Tensor();
|
auto wt8 = masks[0] ? torch::empty({k, n}, fp8_options) : torch::Tensor();
|
||||||
auto xt8 = masks[1] ? torch::empty({k, m}, fp8_options) : torch::Tensor();
|
auto xt8 = masks[1] ? torch::empty({k, m}, fp8_options) : torch::Tensor();
|
||||||
|
// w/x transpose-quantize amax goes to a scratch buffer, NOT amax_g: the
|
||||||
|
// gradient scale must only see the gradient's own max-abs.
|
||||||
|
auto amax_t = torch::zeros({1}, g_c.options().dtype(torch::kFloat32));
|
||||||
|
|
||||||
int64_t block = 256;
|
int64_t block = 256;
|
||||||
quantize_kernel<__nv_fp8_e4m3>
|
quantize_kernel<__nv_fp8_e4m3>
|
||||||
@@ -394,8 +397,8 @@ std::tuple<torch::Tensor, torch::Tensor, torch::Tensor> fp8_linear_backward_scal
|
|||||||
transpose_quantize_kernel<__nv_fp8_e4m3>
|
transpose_quantize_kernel<__nv_fp8_e4m3>
|
||||||
<<<blocks, threads, 0, stream.stream()>>>(
|
<<<blocks, threads, 0, stream.stream()>>>(
|
||||||
reinterpret_cast<const __nv_bfloat16*>(w_c.data_ptr()), swi_ptr,
|
reinterpret_cast<const __nv_bfloat16*>(w_c.data_ptr()), swi_ptr,
|
||||||
reinterpret_cast<__nv_fp8_e4m3*>(wt8.data_ptr()), amax_g_ptr,
|
reinterpret_cast<__nv_fp8_e4m3*>(wt8.data_ptr()),
|
||||||
n, k);
|
amax_t.data_ptr<float>(), n, k);
|
||||||
fp8_gemm_into(g8, wt8, grad_input.reshape({m, k}), m, n, k, sg_ptr,
|
fp8_gemm_into(g8, wt8, grad_input.reshape({m, k}), m, n, k, sg_ptr,
|
||||||
sw_ptr, stream.stream());
|
sw_ptr, stream.stream());
|
||||||
}
|
}
|
||||||
@@ -405,13 +408,13 @@ std::tuple<torch::Tensor, torch::Tensor, torch::Tensor> fp8_linear_backward_scal
|
|||||||
transpose_quantize_kernel<__nv_fp8_e4m3>
|
transpose_quantize_kernel<__nv_fp8_e4m3>
|
||||||
<<<g_blocks, threads, 0, stream.stream()>>>(
|
<<<g_blocks, threads, 0, stream.stream()>>>(
|
||||||
reinterpret_cast<const __nv_bfloat16*>(g_c.data_ptr()), sgi_ptr,
|
reinterpret_cast<const __nv_bfloat16*>(g_c.data_ptr()), sgi_ptr,
|
||||||
reinterpret_cast<__nv_fp8_e4m3*>(gt8.data_ptr()), amax_g_ptr,
|
reinterpret_cast<__nv_fp8_e4m3*>(gt8.data_ptr()),
|
||||||
m, n);
|
amax_t.data_ptr<float>(), m, n);
|
||||||
transpose_quantize_kernel<__nv_fp8_e4m3>
|
transpose_quantize_kernel<__nv_fp8_e4m3>
|
||||||
<<<x_blocks, threads, 0, stream.stream()>>>(
|
<<<x_blocks, threads, 0, stream.stream()>>>(
|
||||||
reinterpret_cast<const __nv_bfloat16*>(x_c.data_ptr()), sxi_ptr,
|
reinterpret_cast<const __nv_bfloat16*>(x_c.data_ptr()), sxi_ptr,
|
||||||
reinterpret_cast<__nv_fp8_e4m3*>(xt8.data_ptr()), amax_g_ptr,
|
reinterpret_cast<__nv_fp8_e4m3*>(xt8.data_ptr()),
|
||||||
m, k);
|
amax_t.data_ptr<float>(), m, k);
|
||||||
fp8_gemm_into(gt8, xt8, grad_weight, n, m, k, sg_ptr, sx_ptr,
|
fp8_gemm_into(gt8, xt8, grad_weight, n, m, k, sg_ptr, sx_ptr,
|
||||||
stream.stream());
|
stream.stream());
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user