Skip to content

Commit 2bc2582

Browse files
committed
[lang] Support vector astype
Signed-off-by: Qiqi Xiao <qiqix@nvidia.com>
1 parent e57575b commit 2bc2582

8 files changed

Lines changed: 50 additions & 15 deletions

File tree

experimental/cuda-lang/src/cuda/lang/_ir/op_impl/vector_impl.py

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@
88
ImplRegistry,
99
require_dtype_spec,
1010
)
11+
from cuda.tile._ir.arithmetic_ops import astype
1112
from cuda.tile._ir.cast_ops import implicit_cast
1213
from cuda.tile._ir.core_ops import bind_method, build_tuple, loosely_typed_const
1314
from cuda.tile._ir.ops import strictly_typed_const
@@ -175,6 +176,16 @@ def vector_with_item_impl(
175176
return vector_with_item(self, index, value)
176177

177178

179+
@impl(getattr, overload=(VectorTy, "astype"))
180+
def getattr_vector_astype(object: Var[VectorTy], name: Var):
181+
return bind_method(object, Vector.astype)
182+
183+
184+
@impl(Vector.astype)
185+
def vector_astype_impl(self: Var[VectorTy], dtype: Var) -> Var[VectorTy]:
186+
return astype(self, require_dtype_spec(dtype))
187+
188+
178189
@impl(operator.getitem, overload=(VectorTy, WILDCARD))
179190
def vector_getitem(object: Var[VectorTy], key: Var[ScalarTy]) -> Var[ScalarTy]:
180191
result_dtype = object.get_type().element_dtype

experimental/cuda-lang/src/cuda/lang/_stub/types.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -143,6 +143,16 @@ def with_item(self, index: int, value: T) -> "Vector[T]":
143143
value: New value.
144144
"""
145145

146+
@stub
147+
def astype(self, dtype: "DType") -> "Vector":
148+
"""Convert each element to ``dtype``.
149+
150+
Returns a new vector of the same length with the given dtype.
151+
152+
Args:
153+
dtype: Target data type of the result vector.
154+
"""
155+
146156

147157
class Pointer(Generic[T]):
148158
"""Typed address into a CUDA memory space with low-level load and store operations."""

experimental/cuda-lang/test/test_vectors.py

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -73,6 +73,20 @@ def kernel(out):
7373
assert out.cpu().tolist() == [1, 1, 1, 4]
7474

7575

76+
def test_astype_on_vector():
77+
@cl.kernel
78+
def kernel(inp, out):
79+
vector = inp.get_base_pointer().load(count=4, alignment=16)
80+
halved = vector.astype(cl.float16)
81+
out.get_base_pointer().store(halved, alignment=8)
82+
83+
values = [1.0, 2.0, 3.0, 4.0]
84+
inp = torch.tensor(values, dtype=torch.float32).cuda()
85+
out = torch.zeros(4, dtype=torch.float16).cuda()
86+
cl.launch(torch.cuda.current_stream(), (1,), (1,), kernel, (inp, out))
87+
assert out.cpu().tolist() == values
88+
89+
7690
@pytest.mark.parametrize('length', (2, 4))
7791
def test_vector_tuple(length):
7892
expect = tuple(range(length))

experimental/cuda-lang/tutorial/fp16_gemm_0.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -89,11 +89,11 @@ def _to_float16_vector(values, base, vsize):
8989
"""Convert one FP32 vector slice to FP16."""
9090
return cl.Vector(
9191
*tuple(
92-
cl.float16(values[base + i])
92+
values[base + i]
9393
for i in cl.static_iter(range(vsize))
9494
),
95-
dtype=cl.float16,
96-
)
95+
dtype=cl.float32,
96+
).astype(cl.float16)
9797

9898

9999
@cl.kernel

experimental/cuda-lang/tutorial/fp16_gemm_1.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -64,11 +64,11 @@ def _to_float16_vector(values, base, vsize):
6464
"""Convert one FP32 vector slice to FP16."""
6565
return cl.Vector(
6666
*tuple(
67-
cl.float16(values[base + i])
67+
values[base + i]
6868
for i in cl.static_iter(range(vsize))
6969
),
70-
dtype=cl.float16,
71-
)
70+
dtype=cl.float32,
71+
).astype(cl.float16)
7272

7373

7474
@cl.kernel

experimental/cuda-lang/tutorial/fp16_gemm_3.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -80,11 +80,11 @@ def _to_float16_vector(values, base, vsize):
8080
"""Convert one FP32 vector slice to FP16."""
8181
return cl.Vector(
8282
*tuple(
83-
cl.float16(values[base + i])
83+
values[base + i]
8484
for i in cl.static_iter(range(vsize))
8585
),
86-
dtype=cl.float16,
87-
)
86+
dtype=cl.float32,
87+
).astype(cl.float16)
8888

8989

9090
@dataclass(frozen=True)

experimental/cuda-lang/tutorial/nvfp4_gemm_1_ws.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -81,11 +81,11 @@ def _to_float16_vector(
8181
) -> cl.Vector[cl.float16]:
8282
return cl.Vector(
8383
*tuple(
84-
cl.float16(values[base + i])
84+
values[base + i]
8585
for i in cl.static_iter(range(count))
8686
),
87-
dtype=cl.float16,
88-
)
87+
dtype=cl.float32,
88+
).astype(cl.float16)
8989

9090

9191
@cl.kernel

experimental/cuda-lang/tutorial/nvfp4_gemm_2_quantize_fp4.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -138,9 +138,9 @@ def _as_float32_vector_64(regs):
138138

139139
def _to_float16_vector(values, base, count):
140140
return cl.Vector(
141-
*tuple(cl.float16(values[base + i]) for i in cl.static_iter(range(count))),
142-
dtype=cl.float16,
143-
)
141+
*tuple(values[base + i] for i in cl.static_iter(range(count))),
142+
dtype=cl.float32,
143+
).astype(cl.float16)
144144

145145

146146
def _pack_fp32x8_to_e2m1x8(values, base):

0 commit comments

Comments
 (0)