Skip to content

Commit 4a0e8f2

Browse files
committed
Update (base update)
[ghstack-poisoned]
1 parent 7013c8d commit 4a0e8f2

85 files changed

Lines changed: 8800 additions & 0 deletions

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

backends/vulkan/custom_ops_lib.py

Lines changed: 122 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -289,6 +289,61 @@ def linear_dq8ca_q4gsw(
289289
lib.impl(name, linear_q4gsw, "CompositeExplicitAutograd")
290290
linear_qc4w_op = getattr(getattr(torch.ops, namespace), name)
291291

292+
293+
# Backward of linear_q4gsw wrt input (for on-device LoRA training through a frozen
294+
# 4-bit base): d_x = d_out @ dequant(W). Reference impl extracts dequant(W) via the
295+
# forward on an identity so it is layout-agnostic; the runtime dispatches this op to
296+
# the tiled q4gsw_backward WGSL kernel (contracts over N).
297+
def linear_q4gsw_backward_impl(
298+
d_out: torch.Tensor,
299+
weights: torch.Tensor,
300+
weight_scales: torch.Tensor,
301+
group_size: int,
302+
) -> torch.Tensor:
303+
in_features = int(weights.shape[1]) * 2
304+
eye = torch.eye(in_features, dtype=d_out.dtype, device=d_out.device)
305+
w_t = linear_q4gsw(eye, weights, weight_scales, group_size) # [in, out]
306+
return d_out @ w_t.t() # [M, out] @ [out, in] = [M, in]
307+
308+
309+
def linear_q4gsw_backward_meta(
310+
d_out: torch.Tensor,
311+
weights: torch.Tensor,
312+
weight_scales: torch.Tensor,
313+
group_size: int,
314+
) -> torch.Tensor:
315+
return d_out.new_empty(d_out.shape[:-1] + (int(weights.shape[1]) * 2,))
316+
317+
318+
name = "linear_q4gsw_backward"
319+
lib.define(
320+
f"{name}(Tensor d_out, Tensor weights, Tensor weight_scales, int group_size) -> Tensor"
321+
)
322+
lib.impl(name, linear_q4gsw_backward_impl, "CompositeExplicitAutograd")
323+
lib.impl(name, linear_q4gsw_backward_meta, "Meta")
324+
linear_q4gsw_backward_op = getattr(getattr(torch.ops, namespace), name)
325+
326+
327+
def linear_q4gsw_setup_context(ctx, inputs, output) -> None:
328+
_x, weights, weight_scales, group_size, _bias = inputs
329+
ctx.save_for_backward(weights, weight_scales)
330+
ctx.group_size = group_size
331+
332+
333+
def linear_q4gsw_backward(ctx, grad_out):
334+
weights, weight_scales = ctx.saved_tensors
335+
d_x = torch.ops.et_vk.linear_q4gsw_backward(
336+
grad_out, weights, weight_scales, ctx.group_size
337+
)
338+
return d_x, None, None, None, None # grads for (x, weights, scales, group_size, bias)
339+
340+
341+
torch.library.register_autograd(
342+
f"{namespace}::linear_q4gsw",
343+
linear_q4gsw_backward,
344+
setup_context=linear_q4gsw_setup_context,
345+
)
346+
292347
name = "linear_dq8ca_q4gsw"
293348
lib.define(
294349
f"""
@@ -1090,3 +1145,70 @@ def rms_norm_impl(
10901145
lib.define(f"{name}(Tensor x, Tensor weight, float eps) -> Tensor")
10911146
lib.impl(name, rms_norm_impl, "CompositeExplicitAutograd")
10921147
rms_norm_op = getattr(getattr(torch.ops, namespace), name)
1148+
1149+
1150+
# STE weight gradient d_out^T @ x through the frozen 4-bit linear_q4gsw base.
1151+
def linear_q4gsw_dw_impl(
1152+
d_out: torch.Tensor,
1153+
x: torch.Tensor,
1154+
) -> torch.Tensor:
1155+
return d_out.reshape(-1, d_out.shape[-1]).t() @ x.reshape(-1, x.shape[-1])
1156+
1157+
1158+
def linear_q4gsw_dw_meta(
1159+
d_out: torch.Tensor,
1160+
x: torch.Tensor,
1161+
) -> torch.Tensor:
1162+
return d_out.new_empty((d_out.shape[-1], x.shape[-1]))
1163+
1164+
1165+
name = "linear_q4gsw_dw"
1166+
lib.define(f"{name}(Tensor d_out, Tensor x) -> Tensor")
1167+
lib.impl(name, linear_q4gsw_dw_impl, "CompositeExplicitAutograd")
1168+
lib.impl(name, linear_q4gsw_dw_meta, "Meta")
1169+
linear_q4gsw_dw_op = getattr(getattr(torch.ops, namespace), name)
1170+
1171+
1172+
1173+
##################
1174+
## q4gsw_requant ##
1175+
##################
1176+
1177+
1178+
# STE re-quant of fp32 latent weights into the frozen-scale 4-bit codes.
1179+
def q4gsw_requant_impl(
1180+
latent: torch.Tensor,
1181+
scales: torch.Tensor,
1182+
group_size: int,
1183+
) -> torch.Tensor:
1184+
n, k = latent.shape
1185+
group_idx = torch.arange(k, device=latent.device) // group_size
1186+
scale_full = scales.t()[:, group_idx] # [N, K]: scales[k // group_size, n]
1187+
nonzero = scale_full != 0
1188+
safe = torch.where(nonzero, scale_full, torch.ones_like(scale_full))
1189+
q = torch.round(latent / safe)
1190+
q = torch.where(nonzero, q, torch.zeros_like(q))
1191+
codes = (torch.clamp(q, -8, 7).to(torch.int32) + 8) & 0xF # [N, K] in 0..15
1192+
k_packed = (k + 1) // 2
1193+
packed = torch.zeros((n, k_packed), dtype=torch.uint8, device=latent.device)
1194+
packed[:, :] = codes[:, 0::2].to(torch.uint8)
1195+
if k > 1:
1196+
high = codes[:, 1::2].to(torch.uint8)
1197+
packed[:, : high.shape[1]] |= high << 4
1198+
return packed
1199+
1200+
1201+
def q4gsw_requant_meta(
1202+
latent: torch.Tensor,
1203+
scales: torch.Tensor,
1204+
group_size: int,
1205+
) -> torch.Tensor:
1206+
n, k = latent.shape
1207+
return latent.new_empty((n, (k + 1) // 2), dtype=torch.uint8)
1208+
1209+
1210+
name = "q4gsw_requant"
1211+
lib.define(f"{name}(Tensor latent, Tensor scales, int group_size) -> Tensor")
1212+
lib.impl(name, q4gsw_requant_impl, "CompositeExplicitAutograd")
1213+
lib.impl(name, q4gsw_requant_meta, "Meta")
1214+
q4gsw_requant_op = getattr(getattr(torch.ops, namespace), name)

backends/vulkan/op_registry.py

Lines changed: 60 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -462,6 +462,15 @@ def register_quantizedlinear_cpp_ops():
462462
)
463463

464464

465+
@update_features(exir_ops.edge.et_vk.linear_q4gsw_backward.default)
466+
def register_linear_q4gsw_backward():
467+
return OpFeatures(
468+
inputs_storage=utils.CONTIGUOUS_ANY,
469+
inputs_dtypes=utils.FP_T,
470+
supports_prepacking=True,
471+
)
472+
473+
465474
@update_features(exir_ops.edge.et_vk.linear_dq8ca_q4gsw.default)
466475
def register_linear_dq8ca_q4gsw():
467476
return OpFeatures(
@@ -1746,6 +1755,57 @@ def register_rms_norm():
17461755
)
17471756

17481757

1758+
1759+
1760+
@update_features(
1761+
[
1762+
exir_ops.edge.aten.ne.Scalar,
1763+
exir_ops.edge.aten.lt.Scalar,
1764+
exir_ops.edge.aten.le.Scalar,
1765+
exir_ops.edge.aten.ge.Scalar,
1766+
]
1767+
)
1768+
def register_compare_scalar_ops():
1769+
return OpFeatures(
1770+
inputs_storage=utils.ANY_STORAGE,
1771+
inputs_dtypes=utils.FP_INT_T,
1772+
outputs_dtypes=utils.BOOL_T,
1773+
supports_resize=True,
1774+
supports_highdim=True,
1775+
)
1776+
1777+
1778+
1779+
1780+
@update_features(exir_ops.edge.aten.logical_not.default)
1781+
def register_logical_not():
1782+
return OpFeatures(
1783+
inputs_storage=utils.ANY_STORAGE,
1784+
inputs_dtypes=utils.BOOL_T,
1785+
supports_resize=True,
1786+
supports_highdim=True,
1787+
)
1788+
1789+
1790+
@update_features(exir_ops.edge.et_vk.linear_q4gsw_dw.default)
1791+
def register_linear_q4gsw_dw():
1792+
return OpFeatures(
1793+
inputs_storage=utils.CONTIGUOUS_ANY,
1794+
inputs_dtypes=utils.FP_T,
1795+
supports_prepacking=True,
1796+
)
1797+
1798+
1799+
1800+
1801+
@update_features(exir_ops.edge.et_vk.q4gsw_requant.default)
1802+
def register_q4gsw_requant():
1803+
return OpFeatures(
1804+
inputs_storage=utils.CONTIGUOUS_ANY,
1805+
inputs_dtypes=utils.FP_T,
1806+
)
1807+
1808+
17491809
#######################
17501810
## Utility functions ##
17511811
#######################

backends/webgpu/CMakeLists.txt

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -38,6 +38,7 @@ set(WEBGPU_SRCS
3838
runtime/ops/sdpa/Sdpa.cpp
3939
runtime/ops/select_as_symint/SelectAsSymint.cpp
4040
runtime/ops/quantized_linear/QuantizedLinear.cpp
41+
runtime/ops/quantized_linear/QuantizedLinearBackward.cpp
4142
runtime/ops/mul/BinaryOp.cpp
4243
runtime/ops/embedding_q4gsw/EmbeddingQ4gsw.cpp
4344
runtime/ops/rope/RotaryEmbedding.cpp
@@ -52,6 +53,24 @@ set(WEBGPU_SRCS
5253
runtime/ops/cat/Cat.cpp
5354
runtime/ops/index/Index.cpp
5455
runtime/ops/sdpa_fd_decode/SdpaFdDecode.cpp
56+
runtime/ops/mm/Mm.cpp
57+
runtime/ops/log_softmax/LogSoftmax.cpp
58+
runtime/ops/softmax/Softmax.cpp
59+
runtime/ops/bmm/Bmm.cpp
60+
runtime/ops/reduce/Reduce.cpp
61+
runtime/ops/div/BinaryOp.cpp
62+
runtime/ops/sub/BinaryOp.cpp
63+
runtime/ops/where/Where.cpp
64+
runtime/ops/compare/Compare.cpp
65+
runtime/ops/gather/Gather.cpp
66+
runtime/ops/expand_copy/ExpandCopy.cpp
67+
runtime/ops/fill/Fill.cpp
68+
runtime/ops/dim_order/DimOrder.cpp
69+
runtime/ops/linear/Linear.cpp
70+
runtime/ops/embedding/Embedding.cpp
71+
runtime/ops/logical_not/LogicalNot.cpp
72+
runtime/ops/quantized_linear/QuantizedLinearDw.cpp
73+
runtime/ops/quantized_linear/QuantizedLinearRequant.cpp
5574
)
5675

5776
add_library(webgpu_backend ${WEBGPU_SRCS})

0 commit comments

Comments
 (0)