Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
39 changes: 33 additions & 6 deletions quadrants/rhi/metal/metal_device.mm
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,8 @@
#include "quadrants/rhi/impl_support.h"
#include "spirv_msl.hpp"

#include <cstdlib>

namespace quadrants::lang {
namespace metal {

Expand Down Expand Up @@ -115,6 +117,14 @@
if (feature_64_bit_integer_math) {
options.set_msl_version(2, 3, 0);
}
// OpAtomicFAddEXT -> SPIRV-Cross emits `atomic_float` / atomic_fetch_add_explicit, which require
// Metal Shading Language 3.0. Without this (and a matching MTLCompileOptions languageVersion in
// get_mtl_library), newLibraryWithSource fails with "unknown type name 'atomic_float'" and the
// subsequent nil computeFunction triggers an ObjC assert (Abort trap: 6).
if (caps.contains(DeviceCapability::spirv_has_atomic_float_add) ||
caps.contains(DeviceCapability::spirv_has_atomic_float)) {
options.set_msl_version(3, 0, 0);
}

compiler.set_msl_options(options);

Expand All @@ -135,8 +145,16 @@
}

MTLLibrary_id mtl_library = device.get_mtl_library(msl);
if (mtl_library == nil) {
return nullptr;
}

MTLFunction_id mtl_function = device.get_mtl_function(mtl_library, std::string("main0"));
if (mtl_function == nil) {
// Avoid -[MTLComputePipelineDescriptorInternal setComputeFunction:]: `computeFunction must not
// be nil` which hard-aborts the process (Abort trap: 6) instead of returning RhiResult::error.
return nullptr;
}

MTLComputePipelineState_id mtl_compute_pipeline_state = nil;
{
Expand Down Expand Up @@ -1086,11 +1104,8 @@ DeviceCapabilityConfig collect_metal_device_caps(MTLDevice_id mtl_device) {
caps.set(DeviceCapability::spirv_has_atomic_int64, 1);
}
if (feature_floating_point_atomics) {
// FIXME: (penguinliong) For some reason floating point atomics doesn't
// work and breaks the FEM99/FEM128 examples. Should consider add them back
// figured out why.
// caps.set(DeviceCapability::spirv_has_atomic_float, 1);
// caps.set(DeviceCapability::spirv_has_atomic_float_add, 1);
caps.set(DeviceCapability::spirv_has_atomic_float, 1);
caps.set(DeviceCapability::spirv_has_atomic_float_add, 1);
Comment on lines +1107 to +1108

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge Gate native float atomics on MSL 3.0 availability

On Apple7/Mac2 hardware running macOS 11/12 (or iOS before 16), this advertises native float atomics even though get_mtl_library() only selects MTLLanguageVersion3_0 inside the macOS 13/iOS 16 availability check. Float-add kernels are consequently emitted with MSL 3.0 atomic_float syntax and rejected by the older runtime instead of using the previously working CAS fallback; gate these capabilities on the same OS availability condition.

Useful? React with 👍 / 👎.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

we are only supporitng mac 14+, https://github.com/Genesis-Embodied-AI/quadrants#installation , so I think this isnt an issue.

Comment on lines +1107 to +1108

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Gate native float atomics on Metal 3 availability

On macOS 11/12 or iOS <16 devices that still report Apple7/Mac2 support, these caps make f32 atomic_add lower to OpAtomicFAddEXT/MSL atomic_float, but get_mtl_library() only sets MTLLanguageVersion3_0 inside @available(macOS 13.0, iOS 16.0, *) and otherwise compiles with the default language version, so those kernels fail instead of using the existing CAS fallback. This regresses installs the project still tags as macosx_11_0_arm64; Apple documents MSL 3.0 as macOS 13+/iOS 16+ at https://developer.apple.com/documentation/metal/mtllanguageversion/version3_0?language=objc.

Useful? React with 👍 / 👎.

}
if (feature_simd_scoped_permute_operations || feature_quad_scoped_permute_operations) {
caps.set(DeviceCapability::spirv_has_subgroup_vote, 1);
Expand Down Expand Up @@ -1501,7 +1516,19 @@ void get_binding_mappings(spirv_cross::SmallVector<spirv_cross::Resource> *resou
MTLLibrary_id mtl_library = nil;
NSError *err = nil;
NSString *msl_ns = [[NSString alloc] initWithUTF8String:source.c_str()];
mtl_library = [mtl_device_ newLibraryWithSource:msl_ns options:nil error:&err];
// Match SPIRV-Cross's MSL version. `atomic_float` (from OpAtomicFAddEXT) is only valid under
// MTLLanguageVersion3_0+; compiling with options:nil rejects it as "unknown type name".
MTLCompileOptions *compile_opts = nil;
DeviceCapabilityConfig caps = get_caps();
if (caps.contains(DeviceCapability::spirv_has_atomic_float_add) ||
caps.contains(DeviceCapability::spirv_has_atomic_float)) {
compile_opts = [[MTLCompileOptions alloc] init];
if (@available(macOS 13.0, iOS 16.0, *)) {
compile_opts.languageVersion = MTLLanguageVersion3_0;
}
}
mtl_library = [mtl_device_ newLibraryWithSource:msl_ns options:compile_opts error:&err];
[compile_opts release];
[msl_ns release];

if (mtl_library == nil) {
Expand Down
178 changes: 178 additions & 0 deletions tests/python/test_fem99_headless.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,178 @@
"""Headless numerical repro for the alleged Metal native-float-atomic FEM99/FEM128 bug.

Background
----------
In Jan 2023 (Taichi #7093, PENGUINLIONG), when Metal switched to SPIR-V codegen, native float atomics were detected for
Apple7+/Mac2+ but immediately commented out with:

FIXME: floating point atomics doesn't work and breaks the FEM99/FEM128 examples.

Those examples were interactive autodiff neo-Hookean soft-body demos (`python/taichi/examples/simulation/fem99.py`,
later removed from Quadrants). They were NEVER turned into a CI test, and the failure mode (wrong numbers? NaN? hang?
visual explosion?) was never written down. Upstream taichi still carries the identical FIXME.

The critical atomic pattern in FEM99 is the scalar energy reduction under autodiff::

U[None] += V[i] * phi_i # parallel over faces; becomes qd.atomic_add(f32)
with qd.ad.Tape(loss=U): ... # reverse scatter also uses float atomics into pos.grad

This file ports that pattern headlessly and checks for the symptoms we can assert without a GUI: finite energy /
positions, no blow-up, and gradients matching a CPU reference on a small case.

QD_WANTED_ARCHS=metal pytest tests/python/test_fem99_headless.py -v
"""

from __future__ import annotations

import numpy as np

import quadrants as qd

from tests import test_utils


def _run_fem99(n_grid: int, n_frames: int, substeps: int, seed: int = 0):
"""Port of the removed fem99.py, headless. Returns (U_hist, pos_final)."""
N = n_grid
dt = 1e-4
dx = 1.0 / N
rho = 4e1
NF = 2 * N**2
NV = (N + 1) ** 2
E, nu = 4e4, 0.2
mu, lam = E / 2 / (1 + nu), E * nu / (1 + nu) / (1 - 2 * nu)
ball_pos = qd.Vector([0.5, 0.0])
ball_radius = 0.32
gravity = qd.Vector([0.0, -40.0])
damping = 12.5

pos = qd.Vector.field(2, float, NV, needs_grad=True)
vel = qd.Vector.field(2, float, NV)
f2v = qd.Vector.field(3, int, NF)
B = qd.Matrix.field(2, 2, float, NF)
F = qd.Matrix.field(2, 2, float, NF, needs_grad=True)
V = qd.field(float, NF)
phi = qd.field(float, NF)
U = qd.field(float, (), needs_grad=True)

@qd.kernel
def update_U():
for i in range(NF):
ia, ib, ic = f2v[i]
a, b, c = pos[ia], pos[ib], pos[ic]
V[i] = abs((a - c).cross(b - c))
D_i = qd.Matrix.cols([a - c, b - c])
F[i] = D_i @ B[i]
for i in range(NF):
F_i = F[i]
log_J_i = qd.log(F_i.determinant())
phi_i = mu / 2 * ((F_i.transpose() @ F_i).trace() - 2)
phi_i -= mu * log_J_i
phi_i += lam / 2 * log_J_i**2
phi[i] = phi_i
# THE atomic float reduction that motivated the Metal native-float-atomic disable.
U[None] += V[i] * phi_i

@qd.kernel
def advance():
for i in range(NV):
acc = -pos.grad[i] / (rho * dx**2)
vel[i] += dt * (acc + gravity)
vel[i] *= qd.exp(-dt * damping)
for i in range(NV):
disp = pos[i] - ball_pos
disp2 = disp.norm_sqr()
if disp2 <= ball_radius**2:
NoV = vel[i].dot(disp)
if NoV < 0:
vel[i] -= NoV * disp / disp2
cond = ((pos[i] < 0) & (vel[i] < 0)) | ((pos[i] > 1) & (vel[i] > 0))
for j in qd.static(range(pos.n)):
if cond[j]:
vel[i][j] = 0
pos[i] += dt * vel[i]

@qd.kernel
def init_pos():
for i, j in qd.ndrange(N + 1, N + 1):
k = i * (N + 1) + j
pos[k] = qd.Vector([i, j]) / N * 0.25 + qd.Vector([0.45, 0.45])
vel[k] = qd.Vector([0.0, 0.0])
for i in range(NF):
ia, ib, ic = f2v[i]
a, b, c = pos[ia], pos[ib], pos[ic]
B_i_inv = qd.Matrix.cols([a - c, b - c])
B[i] = B_i_inv.inverse()

@qd.kernel
def init_mesh():
for i, j in qd.ndrange(N, N):
k = (i * N + j) * 2
a = i * (N + 1) + j
b = a + 1
c = a + N + 2
d = a + N + 1
f2v[k + 0] = [a, b, c]
f2v[k + 1] = [c, d, a]

init_mesh()
init_pos()

u_hist = []
for _ in range(n_frames):
for _ in range(substeps):
with qd.ad.Tape(loss=U):
update_U()
advance()
u_hist.append(float(U[None]))

return np.array(u_hist, dtype=np.float64), pos.to_numpy()


@test_utils.test(arch=[qd.cpu, qd.metal])
def test_fem99_headless_stays_finite():
"""Does the FEM99 autodiff+atomic-reduce pattern stay numerically alive on Metal?"""
# fem99 used N=32; keep it for fidelity on Metal. CPU can take the same size.
n_grid = 32
n_frames = 5
substeps = 30 # same as the original demo's per-frame substep count

u_hist, pos = _run_fem99(n_grid=n_grid, n_frames=n_frames, substeps=substeps)

assert np.isfinite(u_hist).all(), f"energy became non-finite: {u_hist}"
assert np.isfinite(pos).all(), "positions became non-finite"
# Soft body starts in [0.45,0.70]^2-ish; after a few frames under gravity it should stay roughly in the unit square
# (the demo clamps at the walls). Explosion => |pos| >> 10.
assert np.max(np.abs(pos)) < 10.0, f"positions exploded: max|pos|={np.max(np.abs(pos))}"
# Energy should not blow up by many orders of magnitude frame-to-frame.
assert np.max(np.abs(u_hist)) < 1e8, f"energy exploded: {u_hist}"
print(f"FEM99_OK u_hist={u_hist.tolist()} max|pos|={float(np.max(np.abs(pos)))}")


@test_utils.test(arch=[qd.cpu, qd.metal])
def test_ad_scalar_atomic_reduce_matches_closed_form():
"""Smaller, sharper check: the exact `loss[None] += x[i]**2` pattern (test_ad_atomic) on Metal.

If native float atomics corrupt either the forward reduction or the reverse scatter, this fails.
"""
N = 64
x = qd.field(dtype=qd.f32, shape=N, needs_grad=True)
loss = qd.field(dtype=qd.f32, shape=(), needs_grad=True)

@qd.kernel
def func():
for i in x:
loss[None] += x[i] ** 2

for i in range(N):
x[i] = float(i) * 0.1

with qd.ad.Tape(loss):
func()

expected = sum((i * 0.1) ** 2 for i in range(N))
assert loss[None] == test_utils.approx(expected, rel=1e-4)
for i in range(N):
assert x.grad[i] == test_utils.approx(2 * i * 0.1, rel=1e-4)

print(f"AD_REDUCE_OK loss={float(loss[None])}")
Loading