Skip to content

Instantly share code, notes, and snippets.

@JulianKnodt
Created August 14, 2026 08:34
Show Gist options
  • Select an option

  • Save JulianKnodt/4fd53d572d33051acb783946924adb79 to your computer and use it in GitHub Desktop.

Select an option

Save JulianKnodt/4fd53d572d33051acb783946924adb79 to your computer and use it in GitHub Desktop.
Some SDFs in pytorch
import torch
import math
def sphere(p, radius=0.55, dim=-1):
return p.norm(keepdim=True, dim=dim) - radius
def box(p, bound=0, dim=-1):
d = p.abs() - bound
x,y = d.split([1,1], dim=-1)
return d.clamp(min=0.).norm(keepdim=True, dim=-1) + x.maximum(y).clamp(max=0.)
def arc(p, theta:float=0, ra:float=0.3, rb:float=0.2):
if type(theta) == torch.Tensor:
...
else:
theta = torch.full_like(p[..., :1], theta, device=p.device, dtype=p.dtype)
sin = theta.sin()
cos = theta.cos()
sin_cos = torch.cat([sin, cos], dim=-1)
x,y = p.split([1,1], dim=-1)
x = x.abs()
p = torch.cat([x, y], dim=-1)
return torch.where(
(cos * x) > (sin * y),
(p - sin_cos * ra).norm(dim=-1, keepdim=True),
(p.norm(dim=-1,keepdim=True) - ra).abs(),
) - rb
def dot(a,b, dim=-1): return (a * b).sum(dim=dim, keepdim=True)
def length(a, dim): return a.norm(dim=dim, keepdim=True)
def uneven_capsule(p, r1=0.1, r2=0.2, h=0.5, dim=-1):
x,y = p.split([1,1], dim=dim)
p = torch.cat([x.abs(),y], dim=dim)
b = (r1-r2)/h
a = math.sqrt(1.0-b*b);
k = dot(p, torch.tensor([-b,a],device=p.device), dim=dim)
return torch.where(
k < 0.0,
length(p, dim) - r1,
torch.where(
k > a * h,
length(p - torch.tensor([0., h], device=p.device), dim) - r2,
dot(p, torch.tensor([a,b], device=p.device), dim=dim) - r1,
)
)
def pentagon(p, r=0.5, dim=-1):
x,y = p.split([1,1], dim=dim)
p = torch.cat([x.abs(),y], dim=dim)
kx = 0.809016994
ky = 0.587785252
vec2 = lambda a,b: torch.tensor([a,b], device=p.device)
p = p - 2.0 * dot(vec2(-kx,ky),p, dim).clamp(max=0.) * vec2(-kx,ky)
p = p - 2.0 * dot(vec2( kx,ky),p, dim).clamp(max=0.) * vec2( kx,ky)
kz = 0.726542528
p = p - torch.cat([
p[..., 0, None].clamp(min=-r*kz,max=r*kz),
torch.full_like(p[..., 0, None], r),
], dim=dim)
return length(p, dim) * p[..., 1, None].sign()
def eq_tri(p, r=1, dim=-1):
k = math.sqrt(3.0);
x, y = p.split([1,1], dim=dim)
x = x.abs() - r
y = y + r/k
p = torch.where(
x + k*y > 0.,
torch.cat([x-k*y, -k * x - y], dim=dim)/2,
torch.cat([x,y], dim=dim),
)
x,y = p.split([1,1], dim=dim)
x = x - x.clamp(min=-2 * r, max=0.)
p = torch.cat([x,y], dim=dim)
return -p.norm(dim=dim, keepdim=True) * y.sign();
def x(p, w=1, r=0.1, dim=-1):
k = math.sqrt(2.)
w = w / k
p = p.abs()
x,y = p.split([1,1], dim=dim)
p = (p + torch.cat([y, -x], dim=dim)) / k
p = p.abs()
x,_ = p.split([1,1], dim=dim)
diff = torch.cat([torch.full_like(x, w), torch.zeros_like(x)], dim=dim)
pmr = p - r
_,pmry = pmr.split([1,1], dim=dim)
return torch.where(
x > w,
(p - diff).norm(dim=dim, keepdim=True) - r,
# TODO the dimensions here don't make sense:
pmry.clamp(min=0.) - pmr.clamp(max=0.).norm(dim=dim,keepdim=True),
)
def pentagram(p, r:float=0.5):
k1x = 0.809016994; # cos(π/ 5) = ¼(√5+1)
k2x = 0.309016994; # sin(π/10) = ¼(√5-1)
k1y = 0.587785252; # sin(π/ 5) = ¼√(10-2√5)
k2y = 0.951056516; # cos(π/10) = ¼√(10+2√5)
k1z = 0.726542528; # tan(π/ 5) = √(5-2√5)
v1 = torch.tensor([ k1x,-k1y])
v2 = torch.tensor([-k1x,-k1y])
v3 = torch.tensor([ k2x,-k2y])
x, y = p.unbind(dim=-1)
p = torch.stack([x.abs(), y], dim=-1)
p = p - 2*dot(v1, p, dim=-1).clamp(min=0)*v1
p = p - 2*dot(v2, p, dim=-1).clamp(min=0)*v2;
x, y = p.unbind(dim=-1)
x = x.abs()
y = y - r
p = torch.stack([x, y], dim=-1)
length = torch.linalg.norm(p - v3 * dot(p, v3, dim=-1).clamp(min=0, max=k1z * r), dim=-1)
sign = torch.sign(y * k2x + x * k2y)
return (length * sign).unsqueeze(-1)
def cross_2d(a,b,dim=-1):
ax, ay = a.unbind(dim)
bx, by = b.unbind(dim)
return (ax * by - ay * bx).unsqueeze(dim)
def mandelbrot(c, iters=50, dim=-1):
# iterate
z = torch.zeros_like(c)
m2 = torch.zeros_like(c[..., :1])
dz = torch.zeros_like(c)
x = torch.tensor([1.0, 0.0])
for i in range(iters):
valid_mask = m2 < 1024
#Z' -> 2·Z·Z' + 1
zx, zy = z.unbind(dim=dim)
dzx, dzy = dz.unbind(dim=dim)
dz = torch.where(
valid_mask,
2 * torch.stack([
zx * dzx - zy * dzy,
zx * dzy + zy * dzx
], dim=dim) + x,
dz,
)
#Z -> Z² + c
z = torch.where(
valid_mask,
torch.stack([zx * zx - zy * zy, 2 * zx * zy], dim=dim) + c,
z,
)
#assert(z.isfinite().all()), f"{i} {z[~z.isfinite()]}"
#z = vec2( z.x*z.x - z.y*z.y, 2.0*z.x*z.y ) + c;
m2 = torch.where(
valid_mask,
dot(z,z, dim=dim),
m2,
)
#distance
#d(c) = |Z|·log|Z|/|Z'|
dot_zz = dot(z, z, dim=dim).clamp(min=1e-20)
dot_dzdz = dot(dz, dz, dim=dim).clamp(min=1e-20)
d = 0.5 * torch.sqrt(dot_zz/dot_dzdz) * dot_zz.log()
assert(d.isfinite().all()), f"{dot_zz.min().item()} {dot_dzdz.min().item()}"
return d
#float sdCircleWave( in vec2 p, in float tb, in float ra )
#{
# tb = 3.1415927*5.0/6.0*max(tb,0.0001);
# vec2 co = ra*vec2(sin(tb),cos(tb));
# p.x = abs(mod(p.x,co.x*4.0)-co.x*2.0);
# vec2 p1 = p;
# vec2 p2 = vec2(abs(p.x-2.0*co.x),-p.y+2.0*co.y);
# float d1 = ((co.y*p1.x>co.x*p1.y) ? length(p1-co) : abs(length(p1)-ra));
# float d2 = ((co.y*p2.x>co.x*p2.y) ? length(p2-co) : abs(length(p2)-ra));
# return min(d1, d2);
#}
def circle_wave(p, tb:float=3.39, rad:float=0.3):
tb = math.pi * 5.0/6.0 * max(tb, 0.0001)
sc = torch.tensor([math.sin(tb), math.cos(tb)], device=p.device)
co = rad * sc
p1x, p1y = torch.unbind(p, dim=-1)
p1x = torch.abs(torch.remainder(p1x,co[0]*4) - co[0]*2)
p2x = (p1x-2*co[0]).abs()
p2y = -p1y+2*co[1]
p1 = torch.stack([p1x, p1y], dim=-1)
p2 = torch.stack([p2x, p2y], dim=-1)
b = (co[1]*p1x) > (co[0]*p1y)
b = b.unsqueeze(-1)
u = torch.where(b, p1, p2)
v = torch.where(b, p2, p1)
d1 = v.norm(dim=-1, keepdim=True) - rad
d2 = (u - co).norm(dim=-1, keepdim=True)
#d1 = torch.where(
# co[1]*p1x > co[0]*p1y,
# (p1 - co).norm(dim=-1),
# (p1.norm(dim=-1) - rad).abs(),
#)
s = torch.where(b, -d1, d1)
#d2 = torch.where(
# co[1]*p2x>co[0]*p2y,
# (p2 - co).norm(dim=-1),
# (p2.norm(dim=-1) - rad).abs(),
#)
return s.sign() * torch.minimum(d1.abs(), d2)
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment