// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package signal import ( "sourcedock.dev/petrbalvin/tensor/internal/base" "sourcedock.dev/petrbalvin/tensor/internal/core" ) import "math" // Discrete cosine and sine transforms, types I to IV, in the // orthonormal convention. Every transform is one padded inverse FFT: // the input carries a per-sample complex phase, the result is read at // a shifted frequency index and rotated by a per-frequency phase, // taking the real half for the cosine families and the imaginary half // for the sine ones. The orthonormal scaling makes each transform its // own inverse or the transpose of its partner, so the inverse entry // points are aliases: IDCT(x, 2) is DCT(x, 3), IDCT(x, 4) is // DCT(x, 4), and the sine family follows the same rule. // DCT computes the orthonormal discrete cosine transform of type // kind (1 to 4) of the vector x. A kind outside 1 to 4, a rank other // than 1, an empty vector, or type I on fewer than two points is an // error. func DCT(x *core.Array, kind int) (*core.Array, error) { return dctdst(x, kind, true) } // IDCT computes the inverse orthonormal DCT: the transpose partner of // the forward transform of the same kind. func IDCT(x *core.Array, kind int) (*core.Array, error) { switch kind { case 1, 4: return dctdst(x, kind, true) case 2: return dctdst(x, 3, true) case 3: return dctdst(x, 2, true) } return nil, base.Errf("IDCT: kind must be 1 to 4, got %d", kind) } // DST computes the orthonormal discrete sine transform of type kind // (1 to 4) of the vector x. func DST(x *core.Array, kind int) (*core.Array, error) { return dctdst(x, kind, false) } // IDST computes the inverse orthonormal DST. func IDST(x *core.Array, kind int) (*core.Array, error) { switch kind { case 1, 4: return dctdst(x, kind, false) case 2: return dctdst(x, 3, false) case 3: return dctdst(x, 2, false) } return nil, base.Errf("IDST: kind must be 1 to 4, got %d", kind) } // dctdst dispatches the eight transforms onto the shared core. The // table per type and direction: the input phase ramp (half weights // where the convention halves an endpoint), the frequency shift into // the padded spectrum, the per-frequency rotation including the // output-side half weights, the picked half and the normalisation. func dctdst(x *core.Array, kind int, cosine bool) (*core.Array, error) { const name = "DCT/DST" if kind < 1 || kind > 4 { return nil, base.Errf("%s: kind must be 1 to 4, got %d", name, kind) } if x.Dtype() == core.Complex { return nil, base.Errf("%s: complex arrays are not supported", name) } if x.NDim() != 1 || x.Len() == 0 { return nil, base.Errf("%s: the input must be a non-empty vector, got shape %s", name, base.ShapeText(x.Shape())) } n := x.Len() if kind == 1 && n < 2 { return nil, base.Errf("%s: type I needs at least two points, got %d", name, n) } vals := make([]float64, n) for i := range n { vals[i] = x.FloatAt(i) } part := func(z complex128) float64 { return real(z) } if !cosine { part = func(z complex128) float64 { return imag(z) } } // phase weights the samples, shift picks the spectrum index, // outer rotates the picked bin, norm scales the answer. phase := func(j int) complex128 { return 1 } outer := func(k int) complex128 { return 1 } norm := func(k int) float64 { return math.Sqrt(2 / float64(n)) } shift, m := 0, 2*n switch kind { case 1: if cosine { m = 2*n - 2 phase = func(j int) complex128 { if j == 0 || j == n-1 { return math.Sqrt2 / 2 } return 1 } outer = func(k int) complex128 { w := 1.0 if k == 0 || k == n-1 { w = math.Sqrt2 / 2 } return complex(w, 0) } norm = func(int) float64 { return math.Sqrt(2 / float64(n-1)) } } else { // sin(π(j+1)(k+1)/(n+1)) on the length-2(n+1) grid: the // shifted, rotated read carries the whole phase, and the // sine family carries no endpoint weights at all. m = 2*n + 2 outer = func(k int) complex128 { return base.CmplxPolar(1, math.Pi*float64(k+1)/float64(n+1)) } norm = func(int) float64 { return math.Sqrt(2 / float64(n+1)) } shift = 1 } case 2: if cosine { // cos(π(2j+1)k/2n): the frequency rotation carries the // (2j+1) odd shift; the k = 0 row halves. outer = func(k int) complex128 { return base.CmplxPolar(1, math.Pi*float64(k)/float64(2*n)) } norm = func(k int) float64 { if k == 0 { return math.Sqrt(1 / float64(n)) } return math.Sqrt(2 / float64(n)) } } else { // sin(π(2j+1)(k+1)/2n): read the shifted bin. outer = func(k int) complex128 { return base.CmplxPolar(1, math.Pi*float64(k+1)/float64(2*n)) } norm = func(k int) float64 { if k == n-1 { return math.Sqrt(1 / float64(n)) } return math.Sqrt(2 / float64(n)) } shift = 1 } case 3: // The transpose of type 2: cos(π(2k+1)j/2n) and // sin(π(j+1)(2k+1)/2n), with the normalisation on the input // index and none on the output; the transpose of the type-2 // weights lands on the samples. phase = func(j int) complex128 { w := 1.0 if (cosine && j == 0) || (!cosine && j == n-1) { w = math.Sqrt(1 / float64(n)) } else { w = math.Sqrt(2 / float64(n)) } return base.CmplxPolar(w, math.Pi*float64(j)/float64(2*n)) } if cosine { outer = func(int) complex128 { return 1 } } else { outer = func(k int) complex128 { return base.CmplxPolar(1, math.Pi*float64(2*k+1)/float64(2*n)) } } norm = func(int) float64 { return 1 } case 4: phase = func(j int) complex128 { return base.CmplxPolar(1, math.Pi*float64(2*j+1)/float64(4*n)) } outer = func(k int) complex128 { return base.CmplxPolar(1, math.Pi*float64(k)/float64(2*n)) } } z := make([]complex128, m) for j := range n { p := phase(j) z[j] = complex(vals[j]*real(p), vals[j]*imag(p)) } w := paddedIFFT(z) out := make([]float64, n) for k := range n { out[k] = norm(k) * part(w[k+shift]*outer(k)) } return floatsFromArrayMust(out, []int{n}), nil } // paddedIFFT computes the unscaled inverse DFT of z: w[k] = // Σ_j z[j]·e^{+2πijk/m} for m = len(z), by conjugating around the // existing transform. The divide-then-multiply pair reproduces the // exact rounding of the old IFFT-then-rescale path, on one buffer and // one copy fewer. z is overwritten and returned. func paddedIFFT(z []complex128) []complex128 { m := len(z) transform(z, +1) for i := range z { z[i] /= complex(float64(m), 0) z[i] *= complex(float64(m), 0) } return z }