yum-mirror/slang

Making it easier to work with shaders

git clone https://git.yummers.dev/yum-mirror/slang

James Helferty (NVIDIA)Enable metal tests (#8446)8086adc90

master
13.4 KiB544 linesraw
1//TEST:SIMPLE(filecheck=METAL): -stage compute -entry computeMain -target metal
2//TEST:SIMPLE(filecheck=GLSL): -stage compute -entry computeMain -target glsl
3//TEST:SIMPLE(filecheck=GLSL_SPIRV): -stage compute -entry computeMain -target spirv -emit-spirv-via-glsl
4//TEST:SIMPLE(filecheck=SPIR): -stage compute -entry computeMain -target spirv -emit-spirv-directly
5//TEST:SIMPLE(filecheck=HLSL): -stage compute -entry computeMain -target hlsl
6//TEST:SIMPLE(filecheck=CUDA): -stage compute -entry computeMain -target cuda
7//TEST:SIMPLE(filecheck=CPP):  -stage compute -entry computeMain -target cpp
8
9//TEST(compute, vulkan):COMPARE_COMPUTE(filecheck-buffer=BUF):-vk -compute -output-using-type -emit-spirv-via-glsl
10//TEST(compute, vulkan):COMPARE_COMPUTE(filecheck-buffer=BUF):-vk -compute -output-using-type -emit-spirv-directly
11//TEST:SIMPLE(filecheck=METALLIB): -target metallib
12
13//TEST(compute, metal):COMPARE_COMPUTE(filecheck-buffer=BUF):-metal -compute -output-using-type -xslang -DMETAL_COMPUTE
14
15//TEST_INPUT:ubuffer(data=[0 1 -1], stride=4):name=inputBuffer
16RWStructuredBuffer<int> inputBuffer;
17
18//TEST_INPUT: ubuffer(data=[0 0 0 0], stride=4):out,name outputBuffer
19RWStructuredBuffer<int> outputBuffer;
20
21// METALLIB: define void @computeMain
22
23// It is unclear why "nextafter" is not working for Metal.
24#define TEST_WHEN_nextafter_WORKS 0
25
26// NOTE: This test is mainly equality comparisons of math functions results
27// against a precise expected value, but the Metal spec defines a minimum
28// accuracy for many of these math functions such that a range of results may
29// be allowed, presumably corresponding to different generations of hardware.
30// Exact comparisons are preferred here for simplicity's sake, but in cases
31// where, e.g., an M1 may yield a different result from an M4, this test checks
32// that the value falls within the documented range instead.
33
34__generic<T:__BuiltinFloatingPointType>
35bool fuzzyCompare(const T value, const T expected, const T epsilon)
36{
37    return
38        value <= (expected + epsilon) &&
39        value >= (expected - epsilon);
40}
41
42__generic<T:__BuiltinFloatingPointType, let N : int>
43bool fuzzyCompare(const vector<T,N> value, const vector<T,N> expected, const vector<T,N> epsilon)
44{
45    return all(
46        value <= (expected + epsilon) &&
47        value >= (expected - epsilon)
48        );
49}
50
51__generic<T:__BuiltinFloatingPointType>
52bool Test_Scalar()
53{
54    // METAL-LABEL: Test_Scalar
55    const T zero = T(inputBuffer[0]);
56    const T one = T(inputBuffer[1]);
57
58    const int zeroInt = int(inputBuffer[0]);
59
60    const T EPS_E2N13 = T(0.0001220703125); // 2^-13
61
62    T outFloat1, outFloat2;
63    int outInt;
64
65    bool voidResult = true;
66
67    // METAL: sincos(
68    // METAL-NOT: sincos(
69    sincos<T>(zero, outFloat1, outFloat2);
70    voidResult = voidResult && zero == outFloat1 && one == outFloat2;
71
72    return voidResult
73        // METAL: acos(
74        // METALLIB: acos.f32
75        && zero == acos<T>(one)
76
77        // METAL: acosh(
78        // METALLIB: acosh.f32
79        && zero == acosh<T>(one)
80
81        // METAL: asin(
82        // METALLIB: asin.f32
83        && zero == asin<T>(zero)
84
85        // METAL: asinh(
86        // METALLIB: asinh.f32
87        && zero == asinh<T>(zero)
88
89        // METAL: atan(
90        // METALLIB: atan.f32
91        && zero == atan<T>(zero)
92
93        // METAL: atan2(
94        // METALLIB: atan2.f32
95        && zero == atan2<T>(zero, one)
96
97        // METAL: atanh(
98        // METALLIB: atanh.f32
99        && zero == atanh<T>(zero)
100
101        // METAL: ceil(
102        // METALLIB: ceil.f32
103        && zero == ceil<T>(zero)
104
105        // METAL: copysign(
106        // METALLIB: bitcast float
107        && zero == copysign<T>(zero, zero)
108
109        // METAL: cos(
110        // METALLIB: cos.f32
111        && one == cos<T>(zero)
112
113        // METAL: cosh(
114        // METALLIB: cosh.f32
115        && one == cosh<T>(zero)
116
117        // METAL: cospi(
118        // METALLIB: cospi.f32
119        && fuzzyCompare<T>(cospi<T>(zero), one, EPS_E2N13)
120
121        // METAL: divide(
122        // METALLIB: fdiv
123        && zero == divide<T>(zero, one)
124
125        // METAL: exp(
126        // METALLIB: exp.f32
127        && one == exp<T>(zero)
128
129        // METAL: exp2(
130        // METALLIB: exp2.f32
131        && one == exp2<T>(zero)
132
133        // METAL: exp10(
134        // METALLIB: exp10.f32
135        && one == exp10<T>(zero)
136
137        // METAL: fabs(
138        // METALLIB: fabs.f32
139        && zero == fabs<T>(zero)
140
141        // METAL: abs(
142        && zero == abs<T>(zero)
143
144        // METAL: fdim(
145        && zero == fdim<T>(zero, zero)
146
147        // METAL: floor(
148        // METALLIB: floor.f32
149        && zero == floor<T>(zero)
150
151        // METAL: fma(
152        // METALLIB: fma.f32
153        && zero == fma(zero, zero, zero)
154
155        // METAL: fmax(
156        // METALLIB: fmax.f32
157        && zero == fmax<T>(zero, zero)
158
159        // METAL: max(
160        && zero == max<T>(zero, zero)
161
162        // METAL: fmax3(
163        // METALLIB: fmax3.f32
164        && zero == fmax3<T>(zero, zero, zero)
165
166        // METAL: max3(
167        && zero == max3<T>(zero, zero, zero)
168
169        // METAL: fmedian3(
170        // METALLIB: fmedian3.f32
171        && zero == fmedian3<T>(zero, zero, zero)
172
173        // METAL: median3(
174        && zero == median3<T>(zero, zero, zero)
175
176        // METAL: fmin(
177        // METALLIB: fmin.f32
178        && zero == fmin<T>(zero, zero)
179
180        // METAL: min(
181        && zero == min<T>(zero, zero)
182
183        // METAL: fmin3(
184        // METALLIB: fmin3.f32
185        && zero == fmin3<T>(zero, zero, zero)
186
187        // METAL: min3(
188        && zero == min3<T>(zero, zero, zero)
189
190        // METAL-COUNT-2: fmod(
191        // METALLIB-COUNT-2: fmod.f32
192        && zero == fmod<T>(zero, one)
193
194        // METAL: fract(
195        // METALLIB: fract.f32
196        && zero == fract<T>(zero)
197
198        // METAL: frexp(
199        // METALLIB: frexp_float
200        && zero == frexp<T>(zero, outInt) && zeroInt == outInt
201
202        // METAL: ldexp(
203        // METALLIB: ldexp.f32
204        && zero == ldexp<T>(zero, zeroInt)
205
206        // METAL: log(
207        // METALLIB: log.f32
208        && zero == log<T>(one)
209
210        // METAL: log2(
211        // METALLIB: log2.f32
212        && zero == log2<T>(one)
213
214        // METAL: log10(
215        // METALLIB: log10.f32
216        && zero == log10<T>(one)
217
218        // METAL: modf(
219        && zero == modf<T>(zero, outFloat1)
220
221#if TEST_WHEN_nextafter_WORKS
222        // M-ETAL: nextafter(
223        && zero == nextafter<T>(zero, zero)
224#endif
225
226        // METAL: pow(
227        // METALLIB: pow.f32
228        && zero == pow<T>(zero, one)
229
230        // METAL: powr(
231        // METALLIB: powr.f32
232        && zero == powr<T>(zero, one)
233
234        // METAL: rint(
235        // METALLIB: rint.f32
236        && zero == rint<T>(zero)
237
238        // METAL: round(
239        // METALLIB: round.f32
240        && zero == round<T>(zero)
241
242        // METAL: rsqrt(
243        // METALLIB: rsqrt.f32
244        && one == rsqrt<T>(one)
245
246        // METAL: sin(
247        // METALLIB: sin.f32
248        && zero == sin<T>(zero)
249
250        // METAL: sinh(
251        // METALLIB: sinh.f32
252        && zero == sinh<T>(zero)
253
254        // METAL: sinpi(
255        // METALLIB: sinpi.f32
256        && zero == sinpi<T>(zero)
257
258        // METAL: sqrt(
259        // METALLIB: sqrt.f32
260        && zero == sqrt<T>(zero)
261
262        // METAL: tan(
263        // METALLIB: tan.f32
264        && zero == tan<T>(zero)
265
266        // METAL: tanh(
267        // METALLIB: tanh.f32
268        && zero == tanh<T>(zero)
269
270        // METAL: tanpi(
271        // METALLIB: tanpi.f32
272        && zero == tanpi<T>(zero)
273
274        // METAL: trunc(
275        && zero == trunc<T>(zero)
276        ;
277
278    // METALLIB: ret
279}
280
281__generic<T:__BuiltinFloatingPointType, let N : int>
282bool Test_Vector()
283{
284    // METAL-LABEL: Test_Vector_0
285    const vector<T,N> zero = T(inputBuffer[0]);
286    const vector<T,N> one = T(inputBuffer[1]);
287
288    const vector<int,N> zeroInt = int(inputBuffer[0]);
289
290    const vector<T,N> EPS_E2N13 = T(0.0001220703125); // 2^-13
291
292    vector<T,N> outFloat1, outFloat2;
293    vector<int,N> outInt;
294
295    bool voidResult = true;
296
297    // METAL: sincos(
298    // METAL-NOT: sincos(
299    sincos<T>(zero, outFloat1, outFloat2);
300    voidResult = voidResult && zero == outFloat1 && one == outFloat2;
301
302    return voidResult
303        // METAL: acos(
304        // METAL-NOT: acos(
305        && zero == acos<T>(one)
306
307        // METAL: acosh(
308        // METAL-NOT: acosh(
309        && zero == acosh<T>(one)
310
311        // METAL: asin(
312        // METAL-NOT: asin(
313        && zero == asin<T>(zero)
314
315        // METAL: asinh(
316        // METAL-NOT: asinh(
317        && zero == asinh<T>(zero)
318
319        // METAL: atan(
320        // METAL-NOT: atan(
321        && zero == atan<T>(zero)
322
323        // METAL: atan2(
324        // METAL-NOT: atan2(
325        && zero == atan2<T>(zero, one)
326
327        // METAL: atanh(
328        // METAL-NOT: atanh(
329        && zero == atanh<T>(zero)
330
331        // METAL: ceil(
332        // METAL-NOT: ceil(
333        && zero == ceil<T>(zero)
334
335        // METAL: copysign(
336        // METAL-NOT: copysign(
337        && zero == copysign<T>(zero, zero)
338
339        // METAL: cos(
340        // METAL-NOT: cos(
341        && one == cos<T>(zero)
342
343        // METAL: cosh(
344        // METAL-NOT: cosh(
345        && one == cosh<T>(zero)
346
347        // METAL: cospi(
348        // METAL-NOT: cospi(
349        && fuzzyCompare<T,N>(cospi<T>(zero), one, EPS_E2N13)
350
351        // METAL: divide(
352        // METAL-NOT: divide(
353        && zero == divide<T>(zero, one)
354
355        // METAL: exp(
356        // METAL-NOT: exp(
357        && one == exp<T>(zero)
358
359        // METAL: exp2(
360        // METAL-NOT: exp2(
361        && one == exp2<T>(zero)
362
363        // METAL: exp10(
364        // METAL-NOT: exp10(
365        && one == exp10<T>(zero)
366
367        // METAL: fabs(
368        // METAL-NOT: fabs(
369        && zero == fabs<T>(zero)
370
371        // METAL: abs(
372        // METAL-NOT: abs(
373        && zero == abs<T>(zero)
374
375        // METAL: fdim(
376        // METAL-NOT: fdim(
377        && zero == fdim<T>(zero, zero)
378
379        // METAL: floor(
380        // METAL-NOT: floor(
381        && zero == floor<T>(zero)
382
383        // METAL: fma(
384        // METAL-NOT: fma(
385        && zero == fma(zero, zero, zero)
386
387        // METAL: fmax(
388        // METAL-NOT: fmax(
389        && zero == fmax<T>(zero, zero)
390
391        // METAL: max(
392        // METAL-NOT: max(
393        && zero == max<T>(zero, zero)
394
395        // METAL: fmax3(
396        // METAL-NOT: fmax3(
397        && zero == fmax3<T>(zero, zero, zero)
398
399        // METAL: max3(
400        // METAL-NOT: max3(
401        && zero == max3<T>(zero, zero, zero)
402
403        // METAL: fmedian3(
404        // METAL-NOT: fmedian3(
405        && zero == fmedian3<T>(zero, zero, zero)
406
407        // METAL: median3(
408        // METAL-NOT: median3(
409        && zero == median3<T>(zero, zero, zero)
410
411        // METAL: fmin(
412        // METAL-NOT: fmin(
413        && zero == fmin<T>(zero, zero)
414
415        // METAL: min(
416        // METAL-NOT: min(
417        && zero == min<T>(zero, zero)
418
419        // METAL: fmin3(
420        // METAL-NOT: fmin3(
421        && zero == fmin3<T>(zero, zero, zero)
422
423        // METAL: min3(
424        // METAL-NOT: min3(
425        && zero == min3<T>(zero, zero, zero)
426
427        // METAL-COUNT-2: fmod(
428        // METAL-NOT: fmod(
429        && zero == fmod<T>(zero, one)
430
431        // METAL: fract(
432        // METAL-NOT: fract(
433        && zero == fract<T>(zero)
434
435        // METAL: frexp(
436        // METAL-NOT: frexp(
437        && zero == frexp<T>(zero, outInt) && all(zeroInt == outInt)
438
439        // METAL: ldexp(
440        // METAL-NOT: ldexp(
441        && zero == ldexp<T>(zero, zeroInt)
442
443        // METAL: log(
444        // METAL-NOT: log(
445        && zero == log<T>(one)
446
447        // METAL: log2(
448        // METAL-NOT: log2(
449        && zero == log2<T>(one)
450
451        // METAL: log10(
452        // METAL-NOT: log10(
453        && zero == log10<T>(one)
454
455        // METAL: modf(
456        // METAL-NOT: modf(
457        && zero == modf<T>(zero, outFloat1)
458
459#if TEST_WHEN_nextafter_WORKS
460        // M-ETAL: nextafter(
461        // METAL-NOT: nextafter(
462        && zero == nextafter<T>(zero, zero)
463#endif
464
465        // METAL: pow(
466        // METAL-NOT: pow(
467        && zero == pow<T>(zero, one)
468
469        // METAL: powr(
470        // METAL-NOT: powr(
471        && zero == powr<T>(zero, one)
472
473        // METAL: rint(
474        // METAL-NOT: rint(
475        && zero == rint<T>(zero)
476
477        // METAL: round(
478        // METAL-NOT: round(
479        && zero == round<T>(zero)
480
481        // METAL: rsqrt(
482        // METAL-NOT: rsqrt(
483        && one == rsqrt<T>(one)
484
485        // METAL: sin(
486        // METAL-NOT: sin(
487        && zero == sin<T>(zero)
488
489        // METAL: sinh(
490        // METAL-NOT: sinh(
491        && zero == sinh<T>(zero)
492
493        // METAL: sinpi(
494        // METAL-NOT: sinpi(
495        && zero == sinpi<T>(zero)
496
497        // METAL: sqrt(
498        // METAL-NOT: sqrt(
499        && zero == sqrt<T>(zero)
500
501        // METAL: tan(
502        // METAL-NOT: tan(
503        && zero == tan<T>(zero)
504
505        // METAL: tanh(
506        // METAL-NOT: tanh(
507        && zero == tanh<T>(zero)
508
509        // METAL: tanpi(
510        // METAL-NOT: tanpi(
511        && zero == tanpi<T>(zero)
512
513        // METAL: trunc(
514        // METAL-NOT: trunc(
515        && zero == trunc<T>(zero)
516        ;
517
518    // METAL-LABEL: Test_Vector_1
519}
520
521[numthreads(1,1,1)]
522void computeMain()
523{
524    // GLSL: void main(
525    // GLSL_SPIRV: OpEntryPoint
526    // SPIR: OpEntryPoint
527    // HLSL: void computeMain(
528    // CUDA: void computeMain(
529    // CPP: void _computeMain(
530
531    bool result = true
532        && Test_Scalar<float>()
533        && Test_Vector<float, 2>()
534        && Test_Vector<float, 3>()
535        && Test_Vector<float, 4>()
536        && Test_Scalar<half>()
537        && Test_Vector<half, 2>()
538        && Test_Vector<half, 3>()
539        && Test_Vector<half, 4>()
540        ;
541
542    // BUF: 1
543    outputBuffer[0] = int(result);
544}