Skip to content

Commit b45c805

Browse files
chapman20jcopybara-github
authored andcommitted
Add reverse-mode AD support and dimension permutation to qwix hijax.
This change implements Vector-Jacobian Product (VJP) for the hijax conversion primitives. It also introduces a PermuteDims primitive, along with permute_dims and transpose functions for HiQArray, including their VJP implementations. Finally, it adds unit tests to verify the backward passes and permutation operations. PiperOrigin-RevId: 958434873
1 parent 25bb76f commit b45c805

2 files changed

Lines changed: 164 additions & 5 deletions

File tree

qwix/contrib/hijax/convert.py

Lines changed: 88 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -17,8 +17,12 @@
1717

1818
# pyrefly: ignore-errors
1919

20+
import functools
21+
2022
import jax
2123
import jax.experimental.hijax as hjx
24+
import jax.numpy as jnp
25+
import numpy as np
2226
import qwix.contrib.hijax.hiqarray as hq
2327

2428

@@ -74,7 +78,7 @@ def __init__(
7478
self.out_aval = hq.HiQArrayTy(out_qvalue_ty, scale_ty, zero_point_ty)
7579
self.params = dict(
7680
quantize_fn=quantize_fn,
77-
**quantize_kwargs,
81+
quantize_kwargs=quantize_kwargs,
7882
)
7983
# For type checking
8084
self.quantize_fn = quantize_fn
@@ -88,6 +92,16 @@ def expand(self, data, scale, zero_point):
8892
)
8993
return hq.HiQArray(quantized_data, scale, zero_point)
9094

95+
# Reverse mode ad
96+
def vjp_fwd(
97+
self, nzs_in, data: jax.Array, scale: jax.Array, zero_point: jax.Array
98+
):
99+
return self(data, scale, zero_point), None
100+
101+
def vjp_bwd_retval(self, res, g, /):
102+
# Use Straight-Through Estimate (STE)
103+
return (g, None, None)
104+
91105

92106
def to_hiqarray(
93107
data: jax.Array,
@@ -128,6 +142,13 @@ def expand(self, qarray: hq.HiQArray):
128142
)
129143
return dequantized_data
130144

145+
# Reverse mode ad
146+
def vjp_fwd(self, nzs_in, qarray: hq.HiQArray):
147+
return self(qarray), None
148+
149+
def vjp_bwd_retval(self, res, g, /):
150+
return (g,)
151+
131152

132153
def from_hiqarray(
133154
qarray: hq.HiQArray, *, dequantize_fn, **dequantize_kwargs
@@ -138,3 +159,69 @@ def from_hiqarray(
138159
ty, dequantize_fn=dequantize_fn, **dequantize_kwargs
139160
)
140161
return from_qarray_instance(qarray)
162+
163+
164+
class PermuteDims(hjx.VJPHiPrimitive):
165+
"""Hijax primitive for permuting dimensions of a HiQArray."""
166+
167+
def __init__(
168+
self,
169+
in_aval: hq.HiQArrayTy,
170+
axes: tuple[int, ...],
171+
):
172+
self.in_avals = (in_aval,)
173+
self.out_aval = self._permute_dims_aval(in_aval, axes)
174+
self.params = dict(axes=axes)
175+
# For pytype warnings
176+
self.axes = axes
177+
super().__init__()
178+
179+
# Private functions
180+
@staticmethod
181+
def _permute_dims_aval(
182+
in_aval: hq.HiQArrayTy, axes: tuple[int, ...]
183+
) -> hq.HiQArrayTy:
184+
inner_fn = functools.partial(
185+
jax.eval_shape, lambda x: jnp.permute_dims(x, axes=axes)
186+
)
187+
188+
def fn(x):
189+
sds = inner_fn(x)
190+
return jax._src.core._sds_aval_mapping(sds) # pylint: disable=protected-access
191+
192+
lo_avals_permuted = jax.tree_util.tree_map(fn, in_aval.lo_ty())
193+
return hq.HiQArrayTy.raise_ty(lo_avals_permuted)
194+
195+
def expand(self, qarray: hq.HiQArray):
196+
return hq.HiQArray(
197+
jnp.permute_dims(qarray.qvalue, self.axes),
198+
jnp.permute_dims(qarray.scale, self.axes),
199+
(
200+
jnp.permute_dims(qarray.zero_point, self.axes)
201+
if qarray.zero_point is not None
202+
else None
203+
),
204+
)
205+
206+
# Reverse mode ad
207+
def vjp_fwd(self, nzs_in, qarray: hq.HiQArray):
208+
return permute_dims(qarray, self.axes), None
209+
210+
def vjp_bwd_retval(self, res, g, /):
211+
inv_perm = tuple(np.argsort(self.axes))
212+
return (jnp.permute_dims(g, inv_perm),)
213+
214+
215+
def permute_dims(qarray: hq.HiQArray, axes: tuple[int, ...]) -> hq.HiQArray:
216+
ty = jax.typeof(qarray)
217+
permute_axes_instance = PermuteDims(ty, axes)
218+
return permute_axes_instance(qarray)
219+
220+
221+
def transpose(qarray: hq.HiQArray) -> hq.HiQArray:
222+
if qarray.ndim < 2:
223+
raise ValueError(f"Called transpose on HiQArray of shape {qarray.shape}")
224+
225+
s = list(range(qarray.ndim))
226+
new_s = s[:-2] + [s[-1], s[-2]]
227+
return permute_dims(qarray, new_s)

tests/contrib/hijax/hiqarray_test.py

Lines changed: 76 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,5 @@
1+
import functools
2+
13
from absl.testing import absltest
24
from absl.testing import parameterized
35
import jax
@@ -10,11 +12,10 @@
1012
class HiqarrayTest(parameterized.TestCase):
1113

1214
def test_as_hiqarray(self):
13-
x = jnp.ones((16, 32, 64))
14-
scale = jnp.ones((1, 1, 64))
15-
zero_point = jnp.zeros((1, 1, 64))
15+
x = jnp.ones((16, 32))
16+
scale = jnp.ones((16, 1))
1617

17-
xq = convert.as_hiqarray(x, scale, zero_point)
18+
xq = convert.as_hiqarray(x, scale, None)
1819

1920
self.assertIsInstance(xq, hiqarray.HiQArray)
2021

@@ -51,6 +52,77 @@ def dequantize_fn(data, scale, zero_point):
5152
self.assertEqual(y.shape, x.shape)
5253
self.assertEqual(y.dtype, x.dtype)
5354

55+
def test_to_hiqarray_bwd(self):
56+
key = jax.random.key(0)
57+
key1, key2 = jax.random.split(key, 2)
58+
x = jax.random.normal(key1, (16, 32))
59+
scale = jax.random.normal(key2, (16, 1))
60+
61+
def quantize_fn(data, scale, zero_point, qtype):
62+
return qarray.quantize_with_scale_zero_point(
63+
data, qtype, scale, zero_point
64+
).qvalue
65+
66+
fn = functools.partial(
67+
convert.to_hiqarray, quantize_fn=quantize_fn, qtype=jnp.int8
68+
)
69+
70+
_, vjp_fn = jax.vjp(fn, x, scale, None)
71+
g = jnp.ones_like(x)
72+
bwd = vjp_fn(g)
73+
74+
self.assertTrue(jnp.allclose(bwd[0], g))
75+
self.assertLess(jnp.max(jnp.abs(bwd[1])), 1e-6)
76+
self.assertIsNone(bwd[2])
77+
78+
def test_from_hiqarray_bwd(self):
79+
x = jnp.ones((16, 32))
80+
scale = jnp.ones((16, 1))
81+
82+
xq = convert.as_hiqarray(x, scale, None)
83+
84+
def dequantize_fn(data, scale, zero_point):
85+
del zero_point # unused
86+
return qarray.call_with_generic_broadcast(jnp.multiply, data, scale)
87+
88+
fn = functools.partial(convert.from_hiqarray, dequantize_fn=dequantize_fn)
89+
90+
_, vjp_fn = jax.vjp(fn, xq)
91+
g = jnp.ones_like(xq.qvalue)
92+
bwd = vjp_fn(g)
93+
94+
self.assertTrue(jnp.allclose(bwd[0], g))
95+
96+
def test_permute_dims(self):
97+
key = jax.random.key(0)
98+
key1, key2 = jax.random.split(key, 2)
99+
x = jax.random.normal(key1, (16, 32))
100+
scale = jax.random.normal(key2, (16, 1))
101+
102+
xq = convert.as_hiqarray(x, scale, None)
103+
104+
yq = convert.permute_dims(xq, (1, 0))
105+
zq = convert.transpose(xq)
106+
107+
self.assertTrue(jnp.allclose(jnp.transpose(x), yq.qvalue))
108+
self.assertTrue(jnp.allclose(jnp.transpose(scale), yq.scale))
109+
self.assertTrue(jnp.allclose(yq.qvalue, zq.qvalue))
110+
self.assertTrue(jnp.allclose(yq.scale, zq.scale))
111+
112+
def test_permute_dims_bwd(self):
113+
key = jax.random.key(0)
114+
key1, key2, key3 = jax.random.split(key, 3)
115+
x = jax.random.normal(key1, (16, 32))
116+
scale = jax.random.normal(key2, (16, 1))
117+
118+
xq = convert.as_hiqarray(x, scale, None)
119+
120+
_, vjp_fn = jax.vjp(convert.transpose, xq)
121+
g = jax.random.normal(key3, (32, 16))
122+
bwd = vjp_fn(g)
123+
124+
self.assertTrue(jnp.allclose(bwd[0], jnp.transpose(g)))
125+
54126

55127
if __name__ == "__main__":
56128
absltest.main()

0 commit comments

Comments
 (0)