yum-mirror/slang

Making it easier to work with shaders

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

Harsh Aggarwal (NVIDIA)Add Optix Intrinsics Coverage (#8159) (#8310)f5fae0108

master
7.0 KiB248 linesraw
1//TEST:SIMPLE(filecheck=CHECK): -target cuda
2//TEST:SIMPLE(filecheck=CHECK_PTX): -target ptx
3
4//TEST_INPUT: set scene = AccelerationStructure
5uniform RaytracingAccelerationStructure scene;
6
7// Ray query and geometry intrinsics - order independent
8//CHECK-DAG: optixGetHitKind()
9//CHECK-DAG: optixGetInstanceIndex()
10//CHECK-DAG: optixGetInstanceId()
11//CHECK-DAG: optixGetObjectRayDirection()
12//CHECK-DAG: optixGetObjectRayOrigin()
13//CHECK-DAG: optixGetPrimitiveIndex()
14//CHECK-DAG: optixGetRayFlags()
15//CHECK-DAG: optixGetRayTmin()
16//CHECK-DAG: optixGetWorldRayDirection()
17//CHECK-DAG: optixGetWorldRayOrigin()
18
19// HitObject intrinsics - uses generic templated call, not individual _N functions
20//CHECK-DAG: optixHitObjectGetAttribute<CustomAttributes_0>
21//CHECK-DAG: slangOptixHitObjectIsMiss
22
23// Control flow intrinsics
24//CHECK-DAG: optixTerminateRay
25//CHECK-DAG: optixTraverse
26//CHECK-DAG: optixMakeHitObject
27//CHECK-DAG: optixIgnoreIntersection
28
29// Traditional SBT data access
30//CHECK-DAG: optixGetSbtDataPointer
31
32// PTX intrinsics validation - using actual PTX names (order independent)
33//CHECK_PTX-DAG: _optix_get_hit_kind
34//CHECK_PTX-DAG: _optix_read_instance_idx
35//CHECK_PTX-DAG: _optix_read_instance_id  
36//CHECK_PTX-DAG: _optix_get_object_ray_direction_x
37//CHECK_PTX-DAG: _optix_get_object_ray_origin_x
38//CHECK_PTX-DAG: _optix_read_primitive_idx
39//CHECK_PTX-DAG: _optix_get_ray_flags
40//CHECK_PTX-DAG: _optix_get_ray_tmin
41//CHECK_PTX-DAG: _optix_get_world_ray_direction_x
42//CHECK_PTX-DAG: _optix_get_world_ray_origin_x
43//CHECK_PTX-DAG: _optix_terminate_ray
44//CHECK_PTX-DAG: _optix_ignore_intersection
45//CHECK_PTX-DAG: _optix_hitobject_get_traverse_data
46//CHECK_PTX-DAG: _optix_hitobject_make_with_traverse_data
47//CHECK_PTX-DAG: _optix_get_sbt_data_ptr_64
48
49struct CustomAttributes
50{
51    float attr0;
52    uint attr1;
53    float attr2;
54    uint attr3;
55    float attr4;
56    uint attr5;
57    float attr6;
58    uint attr7;
59};
60
61struct IntrinsicData
62{
63    // Ray geometry data
64    uint hitKind;
65    uint instanceIndex;
66    uint instanceID;
67    uint primitiveIndex;
68    float3 objectRayDirection;
69    float3 objectRayOrigin;
70    float3 worldRayDirection;
71    float3 worldRayOrigin;
72    float rayTmin;
73    uint rayFlags;
74    
75    // HitObject attributes
76    float hitObjAttr0;
77    uint hitObjAttr1;
78    float hitObjAttr2;
79    uint hitObjAttr3;
80    float hitObjAttr4;
81    uint hitObjAttr5;
82    float hitObjAttr6;
83    uint hitObjAttr7;
84    
85    // HitObject state
86    bool isMiss;
87};
88
89struct RayPayload
90{
91    IntrinsicData data;
92    bool shouldTerminate;
93};
94
95// Traditional SBT data structure
96struct TraditionalSbtData
97{
98    float multiplier;
99    uint flags;
100    float3 color;
101};
102
103// Add a separate payload for anyhit shader testing
104struct AnyHitPayload
105{
106    bool terminateRequested;
107    uint hitCount;
108};
109
110[shader("closesthit")]
111void closestHitShader(
112    uniform TraditionalSbtData sbtData,  // This triggers traditional optixGetSbtDataPointer
113    inout RayPayload payload, 
114    in CustomAttributes attr)
115{
116    // Test basic ray query intrinsics using correct HLSL names
117    payload.data.hitKind = HitKind();
118    payload.data.instanceIndex = InstanceIndex();
119    payload.data.instanceID = InstanceID();
120    payload.data.primitiveIndex = PrimitiveIndex();
121    payload.data.objectRayDirection = ObjectRayDirection();
122    payload.data.objectRayOrigin = ObjectRayOrigin();
123    payload.data.worldRayDirection = WorldRayDirection();
124    payload.data.worldRayOrigin = WorldRayOrigin();
125    payload.data.rayTmin = RayTMin();
126    payload.data.rayFlags = RayFlags();
127    
128    // Test HitObject operations (using a NOP HitObject for simplicity)
129    HitObject hitObj = HitObject::MakeNop();
130    
131    // Test HitObject attribute access with the correct API
132    CustomAttributes hitObjAttrs = hitObj.GetAttributes<CustomAttributes>();
133    payload.data.hitObjAttr0 = hitObjAttrs.attr0;
134    payload.data.hitObjAttr1 = hitObjAttrs.attr1;
135    payload.data.hitObjAttr2 = hitObjAttrs.attr2;
136    payload.data.hitObjAttr3 = hitObjAttrs.attr3;
137    payload.data.hitObjAttr4 = hitObjAttrs.attr4;
138    payload.data.hitObjAttr5 = hitObjAttrs.attr5;
139    payload.data.hitObjAttr6 = hitObjAttrs.attr6;
140    payload.data.hitObjAttr7 = hitObjAttrs.attr7;
141    
142    // Test HitObject state queries
143    payload.data.isMiss = hitObj.IsMiss();
144    
145    // Test traditional SBT data access - this should generate optixGetSbtDataPointer
146    payload.data.rayTmin *= sbtData.multiplier;
147    if (sbtData.flags > 0)
148    {
149        payload.data.rayTmin += sbtData.color.x + sbtData.color.y + sbtData.color.z;
150    }
151    
152    // Mark test as completed
153    payload.shouldTerminate = true;
154}
155
156[shader("anyhit")]
157void anyHitShader(inout AnyHitPayload payload, in CustomAttributes attr)
158{
159    // Test anyhit-specific intrinsics
160    uint hitKind = HitKind();
161    uint instanceID = InstanceID();
162    float rayT = RayTCurrent();
163    
164    // Count hits for testing
165    payload.hitCount++;
166    
167    // Test termination based on some criteria
168    if (payload.terminateRequested || payload.hitCount > 3)
169    {
170        // Test optixTerminateRay - only available in anyhit shaders
171        AcceptHitAndEndSearch();  // This maps to optixTerminateRay for CUDA
172    }
173    else
174    {
175        // Test ignoring hit to continue traversal
176        IgnoreHit();
177    }
178}
179
180[shader("miss")]
181void missShader(inout RayPayload payload)
182{
183    // Initialize with miss data
184    payload.data.hitKind = 0;
185    payload.data.instanceIndex = ~0u;
186}
187
188[shader("raygeneration")]
189void rayGenShader()
190{
191    uint2 index = DispatchRaysIndex().xy;
192    
193    RayPayload payload;
194    payload.shouldTerminate = (index.x % 2) == 0;
195    
196    RayDesc ray;
197    ray.Origin = float3(index.x, index.y, 0);
198    ray.Direction = float3(0, 0, 1);
199    ray.TMin = 0.001f;
200    ray.TMax = 1000.0f;
201    
202    // Test optixTrace through HitObject TraceRay
203    HitObject hit = HitObject::TraceRay(
204        scene,
205        RAY_FLAG_NONE,
206        0xFF,
207        0,
208        1, 
209        0,
210        ray,
211        payload.data
212    );
213    
214    // Test optixMakeHitObject and optixHitObjectGetTraverseData
215    // Create a custom hit object with specific parameters
216    CustomAttributes testAttrs;
217    testAttrs.attr0 = 1.0f;
218    testAttrs.attr1 = 42;
219    
220    uint hitGroupRecordIndex = 0;
221    uint instanceIndex = index.x;
222    uint geometryIndex = 0;
223    uint primitiveIndex = index.y;
224    uint hitKind = 0;
225    
226    // Test optixMakeHitObject through HitObject::MakeHit
227    HitObject customHit = HitObject::MakeHit(
228        hitGroupRecordIndex,
229        scene,
230        instanceIndex,
231        geometryIndex,
232        primitiveIndex,
233        hitKind,
234        ray,
235        testAttrs
236    );
237    
238    // Test optixHitObjectGetTraverseData - this should be called internally
239    // by the HitObject operations and generate the optix call
240    bool isValidHit = customHit.IsHit();
241    if (isValidHit)
242    {
243        // Access traverse data indirectly through HitObject queries
244        CustomAttributes retrievedAttrs = customHit.GetAttributes<CustomAttributes>();
245        payload.data.hitObjAttr0 = retrievedAttrs.attr0;
246        payload.data.hitObjAttr1 = retrievedAttrs.attr1;
247    }
248}