Skip to content

Commit 3aba645

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

89 files changed

Lines changed: 9452 additions & 0 deletions

File tree

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: 172 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,120 @@ 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+
########################
1151+
## fused_ce (training) ##
1152+
########################
1153+
1154+
1155+
def fused_ce_impl(
1156+
logits: torch.Tensor,
1157+
labels: torch.Tensor,
1158+
n_valid: float,
1159+
) -> tuple[torch.Tensor, torch.Tensor]:
1160+
mask = labels >= 0
1161+
safe = labels.clamp(min=0).long()
1162+
lse = torch.logsumexp(logits, dim=-1)
1163+
picked = logits.gather(-1, safe[:, None]).squeeze(-1)
1164+
loss = torch.where(mask, (lse - picked) / n_valid, torch.zeros_like(lse)).sum()
1165+
softmax = torch.softmax(logits, dim=-1)
1166+
onehot = torch.nn.functional.one_hot(safe, logits.shape[-1]).to(logits.dtype)
1167+
dlogits = torch.where(
1168+
mask[:, None], (softmax - onehot) / n_valid, torch.zeros_like(softmax)
1169+
)
1170+
return loss, dlogits
1171+
1172+
1173+
def fused_ce_meta(
1174+
logits: torch.Tensor,
1175+
labels: torch.Tensor,
1176+
n_valid: float,
1177+
) -> tuple[torch.Tensor, torch.Tensor]:
1178+
return logits.new_empty([]), torch.empty_like(logits)
1179+
1180+
1181+
def fused_ce_setup_context(ctx, inputs, output) -> None:
1182+
ctx.save_for_backward(output[1])
1183+
1184+
1185+
def fused_ce_backward(ctx, grad_loss, grad_dlogits):
1186+
(dlogits,) = ctx.saved_tensors
1187+
return grad_loss * dlogits, None, None
1188+
1189+
1190+
name = "fused_ce"
1191+
lib.define(f"{name}(Tensor logits, Tensor labels, float n_valid) -> (Tensor, Tensor)")
1192+
lib.impl(name, fused_ce_impl, "CompositeExplicitAutograd")
1193+
lib.impl(name, fused_ce_meta, "Meta")
1194+
torch.library.register_autograd(
1195+
f"{namespace}::{name}", fused_ce_backward, setup_context=fused_ce_setup_context
1196+
)
1197+
fused_ce_op = getattr(getattr(torch.ops, namespace), name)
1198+
1199+
1200+
# STE weight gradient d_out^T @ x through the frozen 4-bit linear_q4gsw base.
1201+
def linear_q4gsw_dw_impl(
1202+
d_out: torch.Tensor,
1203+
x: torch.Tensor,
1204+
) -> torch.Tensor:
1205+
return d_out.reshape(-1, d_out.shape[-1]).t() @ x.reshape(-1, x.shape[-1])
1206+
1207+
1208+
def linear_q4gsw_dw_meta(
1209+
d_out: torch.Tensor,
1210+
x: torch.Tensor,
1211+
) -> torch.Tensor:
1212+
return d_out.new_empty((d_out.shape[-1], x.shape[-1]))
1213+
1214+
1215+
name = "linear_q4gsw_dw"
1216+
lib.define(f"{name}(Tensor d_out, Tensor x) -> Tensor")
1217+
lib.impl(name, linear_q4gsw_dw_impl, "CompositeExplicitAutograd")
1218+
lib.impl(name, linear_q4gsw_dw_meta, "Meta")
1219+
linear_q4gsw_dw_op = getattr(getattr(torch.ops, namespace), name)
1220+
1221+
1222+
1223+
##################
1224+
## q4gsw_requant ##
1225+
##################
1226+
1227+
1228+
# STE re-quant of fp32 latent weights into the frozen-scale 4-bit codes.
1229+
def q4gsw_requant_impl(
1230+
latent: torch.Tensor,
1231+
scales: torch.Tensor,
1232+
group_size: int,
1233+
) -> torch.Tensor:
1234+
n, k = latent.shape
1235+
group_idx = torch.arange(k, device=latent.device) // group_size
1236+
scale_full = scales.t()[:, group_idx] # [N, K]: scales[k // group_size, n]
1237+
nonzero = scale_full != 0
1238+
safe = torch.where(nonzero, scale_full, torch.ones_like(scale_full))
1239+
q = torch.round(latent / safe)
1240+
q = torch.where(nonzero, q, torch.zeros_like(q))
1241+
codes = (torch.clamp(q, -8, 7).to(torch.int32) + 8) & 0xF # [N, K] in 0..15
1242+
k_packed = (k + 1) // 2
1243+
packed = torch.zeros((n, k_packed), dtype=torch.uint8, device=latent.device)
1244+
packed[:, :] = codes[:, 0::2].to(torch.uint8)
1245+
if k > 1:
1246+
high = codes[:, 1::2].to(torch.uint8)
1247+
packed[:, : high.shape[1]] |= high << 4
1248+
return packed
1249+
1250+
1251+
def q4gsw_requant_meta(
1252+
latent: torch.Tensor,
1253+
scales: torch.Tensor,
1254+
group_size: int,
1255+
) -> torch.Tensor:
1256+
n, k = latent.shape
1257+
return latent.new_empty((n, (k + 1) // 2), dtype=torch.uint8)
1258+
1259+
1260+
name = "q4gsw_requant"
1261+
lib.define(f"{name}(Tensor latent, Tensor scales, int group_size) -> Tensor")
1262+
lib.impl(name, q4gsw_requant_impl, "CompositeExplicitAutograd")
1263+
lib.impl(name, q4gsw_requant_meta, "Meta")
1264+
q4gsw_requant_op = getattr(getattr(torch.ops, namespace), name)

backends/vulkan/op_registry.py

Lines changed: 74 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,71 @@ def register_rms_norm():
17461755
)
17471756

17481757

1758+
# =============================================================================
1759+
# FusedCe.cpp (training)
1760+
# =============================================================================
1761+
1762+
1763+
@update_features(exir_ops.edge.et_vk.fused_ce.default)
1764+
def register_fused_ce():
1765+
return OpFeatures(
1766+
inputs_storage=utils.CONTIGUOUS_ANY,
1767+
inputs_dtypes=[utils.FP_T, utils.INT_T, utils.NONE_T],
1768+
outputs_dtypes=[utils.FP_T, utils.FP_T],
1769+
)
1770+
1771+
1772+
1773+
1774+
@update_features(
1775+
[
1776+
exir_ops.edge.aten.ne.Scalar,
1777+
exir_ops.edge.aten.lt.Scalar,
1778+
exir_ops.edge.aten.le.Scalar,
1779+
exir_ops.edge.aten.ge.Scalar,
1780+
]
1781+
)
1782+
def register_compare_scalar_ops():
1783+
return OpFeatures(
1784+
inputs_storage=utils.ANY_STORAGE,
1785+
inputs_dtypes=utils.FP_INT_T,
1786+
outputs_dtypes=utils.BOOL_T,
1787+
supports_resize=True,
1788+
supports_highdim=True,
1789+
)
1790+
1791+
1792+
1793+
1794+
@update_features(exir_ops.edge.aten.logical_not.default)
1795+
def register_logical_not():
1796+
return OpFeatures(
1797+
inputs_storage=utils.ANY_STORAGE,
1798+
inputs_dtypes=utils.BOOL_T,
1799+
supports_resize=True,
1800+
supports_highdim=True,
1801+
)
1802+
1803+
1804+
@update_features(exir_ops.edge.et_vk.linear_q4gsw_dw.default)
1805+
def register_linear_q4gsw_dw():
1806+
return OpFeatures(
1807+
inputs_storage=utils.CONTIGUOUS_ANY,
1808+
inputs_dtypes=utils.FP_T,
1809+
supports_prepacking=True,
1810+
)
1811+
1812+
1813+
1814+
1815+
@update_features(exir_ops.edge.et_vk.q4gsw_requant.default)
1816+
def register_q4gsw_requant():
1817+
return OpFeatures(
1818+
inputs_storage=utils.CONTIGUOUS_ANY,
1819+
inputs_dtypes=utils.FP_T,
1820+
)
1821+
1822+
17491823
#######################
17501824
## Utility functions ##
17511825
#######################

backends/webgpu/CMakeLists.txt

Lines changed: 20 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,25 @@ 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/fused_ce/FusedCe.cpp
58+
runtime/ops/log_softmax/LogSoftmax.cpp
59+
runtime/ops/softmax/Softmax.cpp
60+
runtime/ops/bmm/Bmm.cpp
61+
runtime/ops/reduce/Reduce.cpp
62+
runtime/ops/div/BinaryOp.cpp
63+
runtime/ops/sub/BinaryOp.cpp
64+
runtime/ops/where/Where.cpp
65+
runtime/ops/compare/Compare.cpp
66+
runtime/ops/gather/Gather.cpp
67+
runtime/ops/expand_copy/ExpandCopy.cpp
68+
runtime/ops/fill/Fill.cpp
69+
runtime/ops/dim_order/DimOrder.cpp
70+
runtime/ops/linear/Linear.cpp
71+
runtime/ops/embedding/Embedding.cpp
72+
runtime/ops/logical_not/LogicalNot.cpp
73+
runtime/ops/quantized_linear/QuantizedLinearDw.cpp
74+
runtime/ops/quantized_linear/QuantizedLinearRequant.cpp
5575
)
5676

5777
add_library(webgpu_backend ${WEBGPU_SRCS})

0 commit comments

Comments
 (0)