refactor: share the four-corner 3D simplex, a lattice hash and an FBM skeleton
Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5.1
parent
b88adba4a7
commit
7eb7ae3594
+27
-1
@@ -9,7 +9,7 @@ import torch
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Callable, Dict, Any, Optional, Tuple
|
||||
|
||||
from ..utils.color_utils import apply_color_scheme
|
||||
from ..utils.color_utils import apply_color_scheme, hsv_to_rgb, interpolate_colors, COLOR_SCHEMES
|
||||
from ..utils.shape_masks import apply_shape_mask, apply_mask_to_tensor
|
||||
from ..utils.noise_utils import create_coordinate_grid
|
||||
from ..core.params import ShaderParams, get_param_value
|
||||
@@ -140,6 +140,32 @@ class BaseNoiseGenerator(ABC):
|
||||
|
||||
return apply_color_scheme(noise, color_scheme, color_intensity, velocity_field, time)
|
||||
|
||||
@staticmethod
|
||||
def palette_channels(field: torch.Tensor, params: ShaderParams) -> torch.Tensor:
|
||||
"""
|
||||
Map one field [B, 1, H, W] onto the chosen colour scheme, as three channels.
|
||||
|
||||
The palette half of domain_warp's _apply_color_variations, for the
|
||||
generators that draw a scalar field, without the alpha it also returned
|
||||
(the field itself again). Channel 0 comes back the same whether one channel
|
||||
or many were asked for, which is the identity every travel-mode basis rests
|
||||
on; the other two are correlated with it, exactly as domain_warp's are.
|
||||
"""
|
||||
scheme = params.color_scheme
|
||||
intensity = params.color_intensity
|
||||
if scheme in ["none", "0"] or intensity <= 0:
|
||||
return field
|
||||
|
||||
t = (field + 1.0) * 0.5
|
||||
if scheme in COLOR_SCHEMES:
|
||||
r, g, b = interpolate_colors([(s[0], s[1]) for s in COLOR_SCHEMES[scheme]], t, field.device)
|
||||
elif scheme in ("rainbow", "hsv"):
|
||||
r, g, b = hsv_to_rgb(t, torch.full_like(t, 0.8), torch.clamp(t + 0.2, 0.0, 1.0))
|
||||
else:
|
||||
return field
|
||||
colours = torch.cat([r, g, b], dim=1) * 2.0 - 1.0
|
||||
return torch.lerp(field.expand(-1, 3, -1, -1), colours, intensity)
|
||||
|
||||
@staticmethod
|
||||
def apply_common_postprocessing(
|
||||
noise: torch.Tensor,
|
||||
|
||||
+123
@@ -0,0 +1,123 @@
|
||||
"""
|
||||
The skeleton the scalar-field generators share.
|
||||
|
||||
domain_warp draws a field, standardises it, masks it, hands channel 0 to
|
||||
fill_channels and draws the rest through the same closures; five golden fixtures
|
||||
pin that generator to the byte, so its fbm_noise stays where it is. Everything
|
||||
here is for the generators added after it, which have no fixtures to keep: the
|
||||
same skeleton, and an FBM whose 3D path uses the four-corner simplex rather than
|
||||
the corner-0 variant those fixtures froze.
|
||||
|
||||
`seed` may be an int or an int64 tensor of seeds, one per channel. The
|
||||
coordinates then grow a leading axis at the first hash and everything after it
|
||||
broadcasts, which is how fill_channels draws the whole channel axis in one call;
|
||||
see shaders/simplex.py.
|
||||
"""
|
||||
import torch
|
||||
|
||||
from .base import BaseNoiseGenerator
|
||||
from .simplex import simplex_2d, simplex_3d_full
|
||||
from ..utils.shape_masks import apply_shape_mask
|
||||
|
||||
MAX_OCTAVES = 8
|
||||
|
||||
|
||||
def time_is_an_axis(time_axis, time):
|
||||
"""
|
||||
Whether to evaluate the field in 3D with time as the third axis.
|
||||
|
||||
core.shader_noise sets `time_axis` once per clip from the frame count. A
|
||||
direct caller that leaves it unset gets the older test, whether time is
|
||||
nonzero, which is what domain_warp's single-image fixtures pin.
|
||||
"""
|
||||
return (time != 0) if time_axis is None else time_axis
|
||||
|
||||
|
||||
def with_time(p, z):
|
||||
"""`p[..., :2]` with a constant third coordinate."""
|
||||
return torch.cat([p[..., :2], torch.full_like(p[..., :1], float(z))], dim=-1)
|
||||
|
||||
|
||||
def simplex_warp(p, seed, strength, frequency):
|
||||
"""Displace `p` by two low-frequency simplex fields, at most `strength` apart."""
|
||||
if strength <= 0.0:
|
||||
return p
|
||||
q = p * frequency
|
||||
dx = simplex_2d(q, seed + 100)
|
||||
dy = simplex_2d(q, seed + 200)
|
||||
return torch.cat([p[..., 0:1] + dx * strength, p[..., 1:2] + dy * strength], dim=-1)
|
||||
|
||||
|
||||
def fbm(p, octaves, seed, persistence=0.5, lacunarity=2.0, time=None, octave_offset=None):
|
||||
"""
|
||||
Layered simplex noise over `p[..., :2]`.
|
||||
|
||||
With `time` the layers are evaluated in 3D, each a little further along the
|
||||
time axis than the last so they do not move in lockstep. `octave_offset` is a
|
||||
translation of `[dx, dy]` per layer index, so the layers are drawn from
|
||||
different parts of the same field.
|
||||
"""
|
||||
total = None
|
||||
amp, freq, norm = 1.0, 1.0, 0.0
|
||||
for i in range(min(max(int(octaves), 1), MAX_OCTAVES)):
|
||||
q = p[..., :2] * freq
|
||||
if octave_offset is not None:
|
||||
q = q + octave_offset * i
|
||||
if time is None:
|
||||
layer = simplex_2d(q, seed + i)
|
||||
else:
|
||||
layer = simplex_3d_full(with_time(q, time * (0.2 + 0.05 * i)), seed + i)
|
||||
total = amp * layer if total is None else total + amp * layer
|
||||
norm += amp
|
||||
amp *= persistence
|
||||
freq *= lacunarity
|
||||
return total / norm
|
||||
|
||||
|
||||
def standardise(field, seed):
|
||||
"""
|
||||
Zero mean and unit deviation: over the whole draw for one seed, per leading
|
||||
slice for a tensor of seeds, so a batched channel matches the draw it replaces.
|
||||
"""
|
||||
if torch.is_tensor(seed):
|
||||
dims = tuple(range(1, field.dim()))
|
||||
return (field - field.mean(dim=dims, keepdim=True)) / (field.std(dim=dims, keepdim=True) + 1e-8)
|
||||
return (field - field.mean()) / (field.std() + 1e-8)
|
||||
|
||||
|
||||
def draw_channels(field, coords, params, seed, target_channels, contrast=1.0, palette=True):
|
||||
"""
|
||||
Fill the channel axis from `field(seed) -> [B, H, W, 1]` (or `[N, B, H, W, 1]`
|
||||
for a tensor of seeds), finishing each draw the way every generator does:
|
||||
standardised, halved, masked, clamped.
|
||||
|
||||
The shape mask is drawn once and lerped into every channel: it is a spatial
|
||||
shape, and apply_shape_mask reseeds the global RNG. Channel 0 stays on the
|
||||
scalar path, which is what keeps a one-channel draw and the travel-mode bases
|
||||
built from it unchanged; the palette, if any, maps that one field to three.
|
||||
"""
|
||||
shape_type = params.shape_type
|
||||
shape_strength = params.shape_strength
|
||||
mask = None
|
||||
if shape_type not in ["none", "0"] and shape_strength > 0:
|
||||
mask = apply_shape_mask(coords, shape_type, params.time, seed, shape_strength)
|
||||
|
||||
def finish(raw, channel_seed):
|
||||
out = standardise(raw, channel_seed) * (0.5 * contrast)
|
||||
if mask is not None:
|
||||
out = torch.lerp(out, out * mask, shape_strength)
|
||||
return torch.clamp(out, -1.0, 1.0)
|
||||
|
||||
base = finish(field(seed), seed).permute(0, 3, 1, 2)
|
||||
if palette:
|
||||
base = BaseNoiseGenerator.palette_channels(base, params)
|
||||
|
||||
def draw(channel_seed):
|
||||
return finish(field(channel_seed), channel_seed).permute(0, 3, 1, 2)
|
||||
|
||||
def draw_many(channel_seeds):
|
||||
seeds = channel_seeds.reshape(-1, 1, 1, 1, 1)
|
||||
# [N, B, H, W, 1] -> [B, N, H, W]
|
||||
return finish(field(seeds), seeds).squeeze(-1).permute(1, 0, 2, 3)
|
||||
|
||||
return BaseNoiseGenerator.fill_channels(draw, base, target_channels, seed, render_many=draw_many)
|
||||
@@ -36,12 +36,20 @@ match per-slice scalar reductions exactly).
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
F2 = 0.5 * (math.sqrt(3.0) - 1.0)
|
||||
G2 = (3.0 - math.sqrt(3.0)) / 6.0
|
||||
F3 = 1.0 / 3.0
|
||||
G3 = 1.0 / 6.0
|
||||
|
||||
# Gradient table for the four-corner 3D simplex, indexed by hash.
|
||||
SIMPLEX_GRADIENTS = torch.tensor([
|
||||
[1, 1, 0], [-1, 1, 0], [1, -1, 0], [-1, -1, 0],
|
||||
[1, 0, 1], [-1, 0, 1], [1, 0, -1], [-1, 0, -1],
|
||||
[0, 1, 1], [0, -1, 1], [0, 1, -1], [0, -1, -1]
|
||||
], dtype=torch.float32)
|
||||
|
||||
|
||||
def _rotate(x, y, seed):
|
||||
"""
|
||||
@@ -213,3 +221,109 @@ def simplex_3d(p, seed, corners=1):
|
||||
if squeeze:
|
||||
result = result.squeeze(0)
|
||||
return result if result.shape[-1] == 1 else result.unsqueeze(-1)
|
||||
|
||||
|
||||
def lattice_hash(ix, iy, seed, iz=None):
|
||||
"""
|
||||
A uniform value in [0, 1) per integer lattice point, from the hash above.
|
||||
|
||||
`ix`, `iy` and `iz` are int64 tensors; `seed` is an int or an int64 tensor
|
||||
of seeds that broadcasts against them and grows the leading axis the way
|
||||
`simplex_2d` does. The cube alone is linear mod 1013 until it wraps int64,
|
||||
which a small seed and small coordinates never do, and then two cells 29 by
|
||||
10 apart share a value; one xorshift round after it breaks that up.
|
||||
"""
|
||||
h = ix * 1619 + iy * 31337 + _fold(seed, 10000) * 2459
|
||||
if iz is not None:
|
||||
h = h + iz * 6971
|
||||
h = h * h * h
|
||||
h = (h ^ (h >> 21)) * 2654435761
|
||||
return torch.remainder(h, 1013).to(torch.float32) / 1013.0
|
||||
|
||||
|
||||
def simplex_3d_full(coords, seed=0):
|
||||
"""
|
||||
Four-corner 3D simplex noise over `coords[..., :3]`.
|
||||
|
||||
temporal_coherent's own function, moved here unchanged so the generators
|
||||
added after it can share it; its golden fixture pins the generator across
|
||||
the move. `seed` is an int, or an int64 tensor of seeds shaped
|
||||
[N, 1, 1, 1, 1] as fill_channels hands them over.
|
||||
"""
|
||||
dim = coords.shape[-1]
|
||||
device = coords.device
|
||||
if torch.is_tensor(seed):
|
||||
# One seed per channel, always shaped [N,1,1,1]. Everything below works
|
||||
# on coords[..., k], which is [B,H,W] while the coordinates are still
|
||||
# shared and [N,B,H,W] once an earlier step has grown the axis; a rank-4
|
||||
# seed broadcasts correctly against both. Deriving the rank from the
|
||||
# coordinates instead collapses the batch axis in the shared case.
|
||||
seed = seed.reshape(-1, 1, 1, 1)
|
||||
|
||||
gradients = SIMPLEX_GRADIENTS.to(device)
|
||||
|
||||
x = coords[..., 0]
|
||||
y = coords[..., 1]
|
||||
z = coords[..., 2] if dim > 2 else torch.zeros_like(x)
|
||||
|
||||
s = (x + y + z) * F3
|
||||
i = torch.floor(x + s)
|
||||
j = torch.floor(y + s)
|
||||
k = torch.floor(z + s)
|
||||
|
||||
t = (i + j + k) * G3
|
||||
x0 = x - (i - t)
|
||||
y0 = y - (j - t)
|
||||
z0 = z - (k - t)
|
||||
|
||||
# Determine simplex
|
||||
x_ge_y = (x0 >= y0).float()
|
||||
y_ge_z = (y0 >= z0).float()
|
||||
x_ge_z = (x0 >= z0).float()
|
||||
|
||||
i1 = x_ge_y * x_ge_z
|
||||
j1 = (1 - x_ge_y) * y_ge_z
|
||||
k1 = (1 - x_ge_z) * (1 - y_ge_z)
|
||||
|
||||
i2 = x_ge_y + (1 - x_ge_y) * x_ge_z
|
||||
j2 = x_ge_y * (1 - x_ge_z) + (1 - x_ge_y)
|
||||
k2 = (1 - x_ge_z) + x_ge_z * (1 - x_ge_y)
|
||||
|
||||
def grad3d(ix, iy, iz, gx, gy, gz):
|
||||
h = (ix * 1619 + iy * 31337 + iz * 6971 + seed * 2459)
|
||||
h = torch.fmod(h * h * h, 1013)
|
||||
grads = F.embedding(h.long() % 12, gradients)
|
||||
return grads[..., 0] * gx + grads[..., 1] * gy + grads[..., 2] * gz
|
||||
|
||||
noise = torch.zeros_like(x0)
|
||||
|
||||
t0 = 0.6 - x0*x0 - y0*y0 - z0*z0
|
||||
mask0 = (t0 >= 0).float()
|
||||
t0 = t0 * t0
|
||||
noise = noise + mask0 * t0 * t0 * grad3d(i, j, k, x0, y0, z0)
|
||||
|
||||
x1 = x0 - i1 + G3
|
||||
y1 = y0 - j1 + G3
|
||||
z1 = z0 - k1 + G3
|
||||
t1 = 0.6 - x1*x1 - y1*y1 - z1*z1
|
||||
mask1 = (t1 >= 0).float()
|
||||
t1 = t1 * t1
|
||||
noise = noise + mask1 * t1 * t1 * grad3d(i + i1, j + j1, k + k1, x1, y1, z1)
|
||||
|
||||
x2 = x0 - i2 + 2.0 * G3
|
||||
y2 = y0 - j2 + 2.0 * G3
|
||||
z2 = z0 - k2 + 2.0 * G3
|
||||
t2 = 0.6 - x2*x2 - y2*y2 - z2*z2
|
||||
mask2 = (t2 >= 0).float()
|
||||
t2 = t2 * t2
|
||||
noise = noise + mask2 * t2 * t2 * grad3d(i + i2, j + j2, k + k2, x2, y2, z2)
|
||||
|
||||
x3 = x0 - 1.0 + 3.0 * G3
|
||||
y3 = y0 - 1.0 + 3.0 * G3
|
||||
z3 = z0 - 1.0 + 3.0 * G3
|
||||
t3 = 0.6 - x3*x3 - y3*y3 - z3*z3
|
||||
mask3 = (t3 >= 0).float()
|
||||
t3 = t3 * t3
|
||||
noise = noise + mask3 * t3 * t3 * grad3d(i + 1, j + 1, k + 1, x3, y3, z3)
|
||||
|
||||
return (noise * 32.0).unsqueeze(-1)
|
||||
|
||||
@@ -6,7 +6,6 @@ between animation frames by treating time as a proper 4th dimension.
|
||||
"""
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import math
|
||||
import logging
|
||||
from typing import Dict, Any, Optional
|
||||
@@ -17,18 +16,11 @@ from ..utils.shape_masks import apply_shape_mask
|
||||
from ..utils.noise_utils import create_coordinate_grid
|
||||
from ..core.params import ShaderParams, get_param_value
|
||||
from ..core.constants import DEFAULT_CHANNELS
|
||||
from .simplex import simplex_3d_full
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# Precomputed gradients for 3D Simplex noise to avoid runtime branching
|
||||
SIMPLEX_GRADIENTS = torch.tensor([
|
||||
[1, 1, 0], [-1, 1, 0], [1, -1, 0], [-1, -1, 0],
|
||||
[1, 0, 1], [-1, 0, 1], [1, 0, -1], [-1, 0, -1],
|
||||
[0, 1, 1], [0, -1, 1], [0, 1, -1], [0, -1, -1]
|
||||
], dtype=torch.float32)
|
||||
|
||||
|
||||
@shader_generator("temporal_coherent", metadata={"description": "Temporally coherent noise for smooth animations"})
|
||||
class TemporalCoherentNoiseGenerator(BaseNoiseGenerator):
|
||||
"""
|
||||
@@ -255,107 +247,12 @@ class TemporalCoherentNoiseGenerator(BaseNoiseGenerator):
|
||||
@staticmethod
|
||||
def _simplex_3d(coords, seed=0):
|
||||
"""
|
||||
Generate 3D simplex noise.
|
||||
|
||||
Args:
|
||||
coords: Coordinate tensor [B, H, W, 3]
|
||||
seed: Random seed
|
||||
|
||||
Returns:
|
||||
Noise tensor [B, H, W, 1]
|
||||
Four-corner 3D simplex noise, [B, H, W, 3] -> [B, H, W, 1].
|
||||
|
||||
Lives in shaders/simplex.py now, so the generators added later can share
|
||||
it; the golden fixture video_temporal_coherent pins this one across the move.
|
||||
"""
|
||||
dim = coords.shape[-1]
|
||||
device = coords.device
|
||||
if torch.is_tensor(seed):
|
||||
# One seed per channel, always shaped [N,1,1,1]. Everything below works
|
||||
# on coords[..., k], which is [B,H,W] while the coordinates are still
|
||||
# shared and [N,B,H,W] once an earlier step has grown the axis; a rank-4
|
||||
# seed broadcasts correctly against both. Deriving the rank from the
|
||||
# coordinates instead collapses the batch axis in the shared case.
|
||||
seed = seed.reshape(-1, 1, 1, 1)
|
||||
|
||||
# Ensure gradients are on the correct device
|
||||
gradients = SIMPLEX_GRADIENTS.to(device)
|
||||
|
||||
x = coords[..., 0]
|
||||
y = coords[..., 1]
|
||||
z = coords[..., 2] if dim > 2 else torch.zeros_like(x)
|
||||
|
||||
F3 = 1.0 / 3.0
|
||||
G3 = 1.0 / 6.0
|
||||
|
||||
s = (x + y + z) * F3
|
||||
i = torch.floor(x + s)
|
||||
j = torch.floor(y + s)
|
||||
k = torch.floor(z + s)
|
||||
|
||||
t = (i + j + k) * G3
|
||||
x0 = x - (i - t)
|
||||
y0 = y - (j - t)
|
||||
z0 = z - (k - t)
|
||||
|
||||
# Determine simplex
|
||||
x_ge_y = (x0 >= y0).float()
|
||||
y_ge_z = (y0 >= z0).float()
|
||||
x_ge_z = (x0 >= z0).float()
|
||||
|
||||
i1 = x_ge_y * x_ge_z
|
||||
j1 = (1 - x_ge_y) * y_ge_z
|
||||
k1 = (1 - x_ge_z) * (1 - y_ge_z)
|
||||
|
||||
i2 = x_ge_y + (1 - x_ge_y) * x_ge_z
|
||||
j2 = x_ge_y * (1 - x_ge_z) + (1 - x_ge_y)
|
||||
k2 = (1 - x_ge_z) + x_ge_z * (1 - x_ge_y)
|
||||
|
||||
# Optimized gradient calculation using embedding lookup
|
||||
def grad3d_optimized(ix, iy, iz, gx, gy, gz):
|
||||
h = (ix * 1619 + iy * 31337 + iz * 6971 + seed * 2459)
|
||||
h = torch.fmod(h * h * h, 1013)
|
||||
h_int = h.long() % 12
|
||||
|
||||
# Lookup gradients from precomputed table
|
||||
grads = F.embedding(h_int, gradients)
|
||||
|
||||
# Dot product
|
||||
return grads[..., 0] * gx + grads[..., 1] * gy + grads[..., 2] * gz
|
||||
|
||||
noise = torch.zeros_like(x0)
|
||||
|
||||
# Corner 0
|
||||
t0 = 0.6 - x0*x0 - y0*y0 - z0*z0
|
||||
mask0 = (t0 >= 0).float()
|
||||
t0 = t0 * t0
|
||||
noise = noise + mask0 * t0 * t0 * grad3d_optimized(i, j, k, x0, y0, z0)
|
||||
|
||||
# Corner 1
|
||||
x1 = x0 - i1 + G3
|
||||
y1 = y0 - j1 + G3
|
||||
z1 = z0 - k1 + G3
|
||||
t1 = 0.6 - x1*x1 - y1*y1 - z1*z1
|
||||
mask1 = (t1 >= 0).float()
|
||||
t1 = t1 * t1
|
||||
noise = noise + mask1 * t1 * t1 * grad3d_optimized(i + i1, j + j1, k + k1, x1, y1, z1)
|
||||
|
||||
# Corner 2
|
||||
x2 = x0 - i2 + 2.0 * G3
|
||||
y2 = y0 - j2 + 2.0 * G3
|
||||
z2 = z0 - k2 + 2.0 * G3
|
||||
t2 = 0.6 - x2*x2 - y2*y2 - z2*z2
|
||||
mask2 = (t2 >= 0).float()
|
||||
t2 = t2 * t2
|
||||
noise = noise + mask2 * t2 * t2 * grad3d_optimized(i + i2, j + j2, k + k2, x2, y2, z2)
|
||||
|
||||
# Corner 3
|
||||
x3 = x0 - 1.0 + 3.0 * G3
|
||||
y3 = y0 - 1.0 + 3.0 * G3
|
||||
z3 = z0 - 1.0 + 3.0 * G3
|
||||
t3 = 0.6 - x3*x3 - y3*y3 - z3*z3
|
||||
mask3 = (t3 >= 0).float()
|
||||
t3 = t3 * t3
|
||||
noise = noise + mask3 * t3 * t3 * grad3d_optimized(i + 1, j + 1, k + 1, x3, y3, z3)
|
||||
|
||||
result = noise * 32.0
|
||||
return result.unsqueeze(-1)
|
||||
return simplex_3d_full(coords, seed)
|
||||
|
||||
|
||||
# Backward compatibility functions
|
||||
|
||||
@@ -0,0 +1,160 @@
|
||||
"""
|
||||
The primitives the scalar-field generators share: the lattice hash, the
|
||||
four-corner 3D simplex, the FBM over them, and the channel-fill skeleton.
|
||||
|
||||
What matters about all of them is the same thing that matters about
|
||||
simplex_2d: a tensor of seeds must give exactly the fields the seeds would give
|
||||
one at a time, or fill_channels' batched path drifts from the scalar path it
|
||||
replaces.
|
||||
"""
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from snk.core.params import ShaderParams
|
||||
from snk.shaders.base import BaseNoiseGenerator
|
||||
from snk.shaders.fbm import draw_channels, fbm, simplex_warp, standardise, with_time
|
||||
from snk.shaders.simplex import lattice_hash, simplex_3d_full
|
||||
from snk.utils.noise_utils import create_coordinate_grid
|
||||
|
||||
CPU = torch.device("cpu")
|
||||
SEEDS = torch.tensor([8888, 8888 + 6151, 4242], dtype=torch.int64).reshape(3, 1, 1, 1, 1)
|
||||
|
||||
|
||||
def _cells(n=24):
|
||||
ix = torch.arange(n, dtype=torch.int64).view(1, n, 1, 1).expand(1, n, n, 1)
|
||||
iy = torch.arange(n, dtype=torch.int64).view(1, 1, n, 1).expand(1, n, n, 1)
|
||||
return ix, iy
|
||||
|
||||
|
||||
def test_lattice_hash_is_uniform_on_the_unit_interval():
|
||||
ix, iy = _cells(64)
|
||||
values = lattice_hash(ix, iy, 8888)
|
||||
assert values.min() >= 0.0 and values.max() < 1.0
|
||||
assert abs(values.mean().item() - 0.5) < 0.02
|
||||
# The x and y jitters of a cell come from different seeds and must not line up.
|
||||
other = lattice_hash(ix, iy, 8888 + 4099)
|
||||
pairs = torch.stack([values.flatten(), other.flatten()])
|
||||
assert abs(torch.corrcoef(pairs)[0, 1].item()) < 0.1
|
||||
|
||||
|
||||
def test_lattice_hash_has_no_repeat_at_the_cube_hash_collision_offset():
|
||||
"""
|
||||
The hash's linear part repeats every (29, -10) cells mod 1013 for as long as
|
||||
the cube does not wrap int64, which at seed 20 and a 40-cell grid it does not.
|
||||
Cells that far apart would then get the same feature point.
|
||||
"""
|
||||
ix, iy = _cells(64)
|
||||
values = lattice_hash(ix, iy, 20)
|
||||
shifted = lattice_hash(ix + 29, iy - 10, 20)
|
||||
pairs = torch.stack([values.flatten(), shifted.flatten()])
|
||||
assert abs(torch.corrcoef(pairs)[0, 1].item()) < 0.1
|
||||
|
||||
|
||||
def test_lattice_hash_batches_over_a_tensor_of_seeds_exactly():
|
||||
ix, iy = _cells()
|
||||
batched = lattice_hash(ix, iy, SEEDS)
|
||||
assert tuple(batched.shape) == (3, 1, 24, 24, 1)
|
||||
for index, seed in enumerate(SEEDS.flatten().tolist()):
|
||||
assert torch.equal(batched[index], lattice_hash(ix, iy, seed))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("hw", [(16, 16), (22, 39)])
|
||||
def test_the_3d_simplex_batches_over_a_tensor_of_seeds_exactly(hw):
|
||||
p = with_time(create_coordinate_grid(1, *hw, CPU) * 3.0, 0.4)
|
||||
batched = simplex_3d_full(p, SEEDS)
|
||||
assert tuple(batched.shape) == (3, 1) + hw + (1,)
|
||||
for index, seed in enumerate(SEEDS.flatten().tolist()):
|
||||
assert torch.equal(batched[index], simplex_3d_full(p, seed))
|
||||
|
||||
|
||||
def test_the_3d_simplex_is_the_function_temporal_coherent_carries():
|
||||
from snk.shaders.temporal_coherent_noise import TemporalCoherentNoiseGenerator
|
||||
|
||||
p = with_time(create_coordinate_grid(1, 16, 16, CPU) * 3.0, 0.4)
|
||||
assert torch.equal(TemporalCoherentNoiseGenerator._simplex_3d(p, 8888), simplex_3d_full(p, 8888))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("time", [None, 0.3])
|
||||
def test_fbm_batches_over_a_tensor_of_seeds_exactly(time):
|
||||
p = create_coordinate_grid(1, 22, 39, CPU) * 2.0
|
||||
batched = fbm(p, 3, SEEDS, time=time)
|
||||
assert tuple(batched.shape) == (3, 1, 22, 39, 1)
|
||||
for index, seed in enumerate(SEEDS.flatten().tolist()):
|
||||
assert torch.equal(batched[index], fbm(p, 3, seed, time=time))
|
||||
|
||||
|
||||
def test_fbm_adds_detail_with_octaves_and_is_bounded():
|
||||
p = create_coordinate_grid(1, 48, 48, CPU) * 2.0
|
||||
|
||||
def neighbour_correlation(field):
|
||||
pairs = torch.stack([field[0, :-1].flatten(), field[0, 1:].flatten()])
|
||||
return torch.corrcoef(pairs)[0, 1].item()
|
||||
|
||||
one = fbm(p, 1, 8888)
|
||||
four = fbm(p, 4, 8888)
|
||||
assert one.abs().max() <= 1.0 and four.abs().max() <= 1.0
|
||||
assert neighbour_correlation(one) > neighbour_correlation(four)
|
||||
|
||||
|
||||
def test_simplex_warp_moves_points_by_at_most_its_strength():
|
||||
p = create_coordinate_grid(1, 16, 16, CPU)
|
||||
assert torch.equal(simplex_warp(p, 8888, 0.0, 0.5), p)
|
||||
warped = simplex_warp(p, 8888, 0.3, 0.5)
|
||||
assert (warped - p).abs().max() <= 0.3
|
||||
assert not torch.equal(warped, p)
|
||||
|
||||
|
||||
def test_standardise_treats_each_seed_slice_on_its_own():
|
||||
field = torch.randn(3, 1, 8, 8, 1) * torch.tensor([1.0, 3.0, 0.2]).view(3, 1, 1, 1, 1) + 2.0
|
||||
per_slice = standardise(field, SEEDS)
|
||||
for index in range(3):
|
||||
assert torch.equal(per_slice[index], standardise(field[index], 8888))
|
||||
assert abs(per_slice[index].mean().item()) < 1e-5
|
||||
assert abs(per_slice[index].std().item() - 1.0) < 1e-3
|
||||
|
||||
|
||||
def _params(**overrides):
|
||||
return ShaderParams({
|
||||
"scale": 1.0, "octaves": 2.0, "warp_strength": 0.5, "phase_shift": 0.5,
|
||||
"shape_type": "none", "shape_strength": 1.0, "color_scheme": "none",
|
||||
"color_intensity": 0.8, "time": 0.0, "base_seed": 8888, **overrides,
|
||||
}).validate()
|
||||
|
||||
|
||||
def _field(coords):
|
||||
def field(seed):
|
||||
return fbm(coords * 2.0, 2, seed)
|
||||
return field
|
||||
|
||||
|
||||
@pytest.mark.parametrize("shape_type", ["none", "radial"])
|
||||
def test_draw_channels_keeps_channel_zero_on_the_scalar_path(shape_type):
|
||||
coords = create_coordinate_grid(1, 22, 39, CPU)
|
||||
params = _params(shape_type=shape_type)
|
||||
wide = draw_channels(_field(coords), coords, params, 8888, 24)
|
||||
single = draw_channels(_field(coords), coords, params, 8888, 1)
|
||||
assert tuple(wide.shape) == (1, 24, 22, 39)
|
||||
assert torch.equal(wide[:, :1], single)
|
||||
assert wide.abs().max() <= 1.0
|
||||
# Every other channel is the draw its own seed would give, drawn alone.
|
||||
for channel in (1, 7):
|
||||
alone = draw_channels(_field(coords), coords, params, 8888 + 6151 * channel, 1)
|
||||
assert torch.equal(wide[:, channel:channel + 1], alone)
|
||||
|
||||
|
||||
def test_draw_channels_does_not_touch_the_global_rng(monkeypatch):
|
||||
coords = create_coordinate_grid(1, 16, 16, CPU)
|
||||
monkeypatch.setattr(torch, "manual_seed", lambda *a, **k: pytest.fail("reseeded the global RNG"))
|
||||
draw_channels(_field(coords), coords, _params(), 8888, 8)
|
||||
|
||||
|
||||
def test_palette_channels_maps_one_field_to_three():
|
||||
field = torch.linspace(-1.0, 1.0, 64).view(1, 1, 8, 8)
|
||||
assert torch.equal(BaseNoiseGenerator.palette_channels(field, _params()), field)
|
||||
coloured = BaseNoiseGenerator.palette_channels(field, _params(color_scheme="viridis"))
|
||||
assert tuple(coloured.shape) == (1, 3, 8, 8)
|
||||
assert coloured.abs().max() <= 1.0
|
||||
# The three are a palette of one field, so they are not copies of each other.
|
||||
assert not torch.equal(coloured[:, 0], coloured[:, 1])
|
||||
rainbow = BaseNoiseGenerator.palette_channels(field, _params(color_scheme="rainbow"))
|
||||
assert tuple(rainbow.shape) == (1, 3, 8, 8)
|
||||
Reference in New Issue
Block a user