// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: MIT package integrate import ( "testing" "sourcedock.dev/petrbalvin/tensor/internal/base" "sourcedock.dev/petrbalvin/tensor/internal/core" ) // The ADI sweeps enforce the boundary constants on the working state's // ring rows so the stencils read neighbours unconditionally, and the // y half-step builds its right sides in column pairs sharing the star // loads. The published history must not see any of that machinery: // sample 0 is the initial state exactly, and every later sample matches // a serial per-line reference bit for bit, non-zero boundary constants // included (the previous pins all carried zero boundaries, so a swapped // constant passed silently). // heat2DReference walks the documented alternating-direction scheme one // line at a time with fresh scratch: the same expressions in the same // order the kernel's lanes use, so equal reads give equal bits and the // comparison below is exact. func heat2DReference(u0 []float64, rows, cols int, kappa, dx, dy, tFinal, dt float64, samples int, bb, bt, bl, br float64) [][]float64 { steps, h := pdeSchedule(tFinal, dt, samples) rx := kappa * h / (2 * dx * dx) ry := kappa * h / (2 * dy * dy) u := append([]float64(nil), u0...) history := [][]float64{append([]float64(nil), u...)} for c := range cols { u[c] = bb u[(rows-1)*cols+c] = bt } every := steps / (samples - 1) star := make([]float64, rows*cols) lowerX := make([]float64, cols-3) upperX := make([]float64, cols-3) diagX := make([]float64, cols-2) for i := range lowerX { lowerX[i] = -rx upperX[i] = -rx } for i := range diagX { diagX[i] = 1 + 2*rx } lowerY := make([]float64, rows-3) upperY := make([]float64, rows-3) diagY := make([]float64, rows-2) for i := range lowerY { lowerY[i] = -ry upperY[i] = -ry } for i := range diagY { diagY[i] = 1 + 2*ry } for s := 1; s <= steps; s++ { clear(star) for r := 1; r < rows-1; r++ { row := u[r*cols : (r+1)*cols] up := u[(r+1)*cols : (r+2)*cols] down := u[(r-1)*cols : r*cols] rhs := make([]float64, cols) for c := range cols { rhs[c] = row[c] + ry*(up[c]-2*row[c]+down[c]) } rhs[1] += rx * bl rhs[cols-2] += rx * br dst := star[r*cols+1 : r*cols+cols-1] err := base.TriSolve(dst, make([]float64, cols-2), make([]float64, cols-2), lowerX, diagX, upperX, rhs[1:cols-1]) if err != nil { panic(err) } star[r*cols] = bl star[r*cols+cols-1] = br } for c := 1; c < cols-1; c++ { rhs := make([]float64, rows) for r := range rows { off := r * cols l := star[off+c-1] wm := star[off+c] e := star[off+c+1] rhs[r] = wm + rx*(e-2*wm+l) } rhs[1] += ry * bb rhs[rows-2] += ry * bt dst := make([]float64, rows-2) err := base.TriSolve(dst, make([]float64, rows-2), make([]float64, rows-2), lowerY, diagY, upperY, rhs[1:rows-1]) if err != nil { panic(err) } u[c] = bb u[(rows-1)*cols+c] = bt for r := 1; r < rows-1; r++ { u[r*cols+c] = dst[r-1] } } for r := range rows { u[r*cols] = bl u[r*cols+cols-1] = br } if s%every == 0 && len(history) < samples { history = append(history, append([]float64(nil), u...)) } } history[samples-1] = append([]float64(nil), u...) return history } func TestHeat2DSampleZeroIsInitialState(t *testing.T) { // The ring enforcement belongs to the working state: sample 0 is // the initial state exactly, boundary constants and corners // included. const rows, cols = 3, 4 u0 := []float64{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12} state, err := core.FromFloats(u0, rows, cols) if err != nil { t.Fatalf("FromFloats: %v", err) } hist, err := IntegrateHeat2D(state, 1, 1, 1, 0.1, 0.05, 2, 1, 2, 3, 4) if err != nil { t.Fatalf("IntegrateHeat2D: %v", err) } got := hist.RawFloats() for i := range u0 { if got[i] != u0[i] { t.Fatalf("sample 0 element %d = %v, want the initial %v", i, got[i], u0[i]) } } } func TestHeat2DSamplesMatchSerialReference(t *testing.T) { const rows, cols = 5, 7 u0 := make([]float64, rows*cols) for i := range u0 { u0[i] = float64(i%13) - 6 } state, err := core.FromFloats(u0, rows, cols) if err != nil { t.Fatalf("FromFloats: %v", err) } const ( kappa, dx, dy = 0.7, 0.3, 0.25 tFinal, dt = 0.08, 0.01 samples = 4 bb, bt, bl, br = 1.5, -2.25, 3.125, -4.5 ) hist, err := IntegrateHeat2D(state, kappa, dx, dy, tFinal, dt, samples, bb, bt, bl, br) if err != nil { t.Fatalf("IntegrateHeat2D: %v", err) } want := heat2DReference(u0, rows, cols, kappa, dx, dy, tFinal, dt, samples, bb, bt, bl, br) got := hist.RawFloats() for s := range samples { for i := range u0 { if got[s*rows*cols+i] != want[s][i] { t.Fatalf("sample %d element %d = %v, want %v", s, i, got[s*rows*cols+i], want[s][i]) } } } }