Files
tensor/signal/dct.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

210 lines
6.5 KiB
Go

// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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
}