210 lines
6.5 KiB
Go
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
|
|
}
|