import jax.numpy as jnp
from jaxcad.backends.wgsl import compile_sdf_to_wgsl
from jaxcad.sdf.boolean import Difference, Union
from jaxcad.sdf.primitives import Box, Capsule, Sphere, Torus
from jaxcad.sdf.transforms import Translate
from jaxcad.viewer import SDFViewerInteractive WebGPU Viewer
jaxcad compiles any SDF scene to WebGPU Shading Language (WGSL) via JAX’s stable HLO export and renders it live in the browser with SDFViewer.
Pipeline:
JAX scene → jax.export → StableHLO MLIR → WGSL → WebGPU fragment shader
Controls — left-drag: orbit · scroll: zoom · right-drag: pan · double-click: reset
1. Sphere — Hello World
SDFViewer takes any jaxcad SDF, compiles it to WGSL, and embeds the result in an interactive WebGPU canvas. Requires Chrome 113+ or Edge 113+.
SDFViewer(Sphere(1.0))2. Boolean operations
sphere = Sphere(1.0)
box = Translate(Box([1.3, 1.3, 1.3]), offset=jnp.array([0.5, 0.3, 0.0]))
scene = Union(sphere, box, smoothness=0.2)
SDFViewer(scene)SDFViewer(Difference(sphere, box, smoothness=0.05))3. Complex scene — blob of spheres
The entire tree is compiled to a single WGSL function — the GPU evaluates the SDF directly with no intermediate mesh.
blob = Union(
Sphere(0.80),
Translate(Sphere(0.70), offset=jnp.array([1.2, 0.2, 0.0])),
Translate(Sphere(0.55), offset=jnp.array([0.5, 0.9, -0.7])),
Translate(Sphere(0.50), offset=jnp.array([-1.0, 0.4, 0.5])),
Translate(Sphere(0.40), offset=jnp.array([0.1, 1.3, 0.3])),
smoothness=0.30,
)
cutter = Translate(Box([0.9, 0.28, 0.9]), offset=jnp.array([0.9, 1.0, 0.4]))
scene = Difference(blob, cutter, smoothness=0.05)
SDFViewer(scene, height=400)4. Torus
Demonstrating a primitive whose SDF uses trigonometric operations — all StableHLO ops compile transparently to WGSL.
scene = Union(
Torus(major_radius=0.9, minor_radius=0.25),
Translate(Sphere(0.3), offset=jnp.array([0.0, 0.6, 0.0])),
smoothness=0.1,
)
SDFViewer(scene)5. Hot-reload — edit the scene without losing camera state
Create a viewer once, then call update_scene() to swap in a new SDF. The camera stays where it is; only the shader recompiles.
viewer = SDFViewer(Sphere(1.0))
viewer# Run this cell to swap the scene — the viewer above updates instantly
viewer.update_scene(
Union(
Sphere(0.9),
Translate(Capsule(radius=0.3, height=1.4), offset=jnp.array([1.2, 0.0, 0.0])),
smoothness=0.15,
)
)6. Inspect the compiled WGSL
compile_sdf_to_wgsl returns the raw shader source — useful for debugging or exporting to a standalone WebGPU application.
print(compile_sdf_to_wgsl(Sphere(1.5)))fn norm(_arg0: vec3<f32>) -> f32 {
let _v0: f32 = 0.000000;
let _v1: vec3<f32> = _arg0 * _arg0;
let _v2: f32 = dot(_v1, vec3<f32>(1.0, 1.0, 1.0));
let _v3: f32 = sqrt(_v2);
return _v3;
}
fn sdf(p: vec3<f32>) -> f32 {
let _v0: f32 = 1.500000;
let _v1: f32 = norm(p);
let _v2: f32 = f32(_v0);
let _v3: f32 = _v1 - _v2;
return _v3;
}
scene = Union(
Sphere(1.0), Translate(Box([0.5, 0.5, 0.5]), offset=jnp.array([1.2, 0.0, 0.0])), smoothness=0.1
)
print(compile_sdf_to_wgsl(scene))fn norm(_arg0: vec3<f32>) -> f32 {
let _v0: f32 = 0.000000;
let _v1: vec3<f32> = _arg0 * _arg0;
let _v2: f32 = dot(_v1, vec3<f32>(1.0, 1.0, 1.0));
let _v3: f32 = sqrt(_v2);
return _v3;
}
fn sdf(p: vec3<f32>) -> f32 {
let _v0: f32 = 0.250000;
let _v1: f32 = 0.000000;
let _v2: f32 = 4.000000;
let _v3: f32 = -1e38;
let _v4: f32 = 0.000000;
let _v5: f32 = 0.000000;
let _v6: f32 = 1.000000;
let _v7: vec3<f32> = vec3<f32>(1.200000, 0.000000, 0.000000);
let _v8: vec3<f32> = vec3<f32>(0.500000, 0.500000, 0.500000);
let _v9: f32 = 0.100000;
let _v10: f32 = norm(p);
let _v11: f32 = f32(_v6);
let _v12: f32 = _v10 - _v11;
let _v13: vec3<f32> = p - _v7;
let _v14: vec3<f32> = abs(_v13);
let _v15: vec3<f32> = _v14 - _v8;
let _v16: vec3<f32> = vec3<f32>(_v5);
let _v17: vec3<f32> = max(_v15, _v16);
let _v18: vec3<f32> = _v17 * _v17;
let _v19: f32 = dot(_v18, vec3<f32>(1.0, 1.0, 1.0));
let _v20: f32 = _v19 + _v4;
let _v21: f32 = sqrt(_v20);
let _v22: f32 = max(max(_v15.x, _v15.y), _v15.z);
let _v23: f32 = min(_v22, _v5);
let _v24: f32 = _v21 + _v23;
let _v25: f32 = _v9 * _v2;
let _v26: f32 = max(_v25, _v1);
let _v27: f32 = _v12 - _v24;
let _v28: f32 = abs(_v27);
let _v29: f32 = f32(_v26);
let _v30: f32 = _v29 - _v28;
let _v31: f32 = max(_v30, _v5);
let _v32: f32 = min(_v12, _v24);
let _v33: f32 = _v31 * _v31;
let _v34: f32 = _v33 * _v0;
let _v35: f32 = f32(_v26);
let _v36: f32 = _v34 / _v35;
let _v37: f32 = _v32 - _v36;
return _v37;
}