Files
rdplib/plugin/rdpgfx/rfx_dwt_shared.go
T

222 lines
8.0 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package rdpgfx
// Standard RemoteFX (MS-RDPRFX) codec helpers shared with rfx.go.
// Extracted from the progressive decoder rewrite; algorithms unchanged.
func rfxGetQuant(quants []rfxQuant, idx int) rfxQuant {
if idx < len(quants) {
return quants[idx]
}
return rfxQuant{6, 6, 6, 6, 6, 6, 6, 6, 6, 6}
}
// rfxDecodeComponent decodes one color component (Y, Cb, or Cr) for a 64×64 tile.
// The returned slice is backed by a *coeffArr from coeffPool; the caller must
// return it via coeffPool.Put((*coeffArr)(result)) when done.
func rfxDecodeComponent(data []byte, quant rfxQuant, rlgrMode int) []int16 {
const tilePixels = rfxTileSize * rfxTileSize // 4096
// Get a pooled coefficient buffer. The pool stores *coeffArr (pointer to a
// fixed-size array) so the any interface stores a single pointer word with no
// heap-boxing allocation.
arr := coeffPool.Get().(*coeffArr)
coeffs := arr[:]
if data == nil {
clear(coeffs)
return coeffs
}
// 1. RLGR entropy decode → 4096 coefficients
if rlgrMode == 3 {
coeffs = rlgr3Decode(data, tilePixels, coeffs)
} else {
coeffs = rlgr1Decode(data, tilePixels, coeffs)
}
// 2. Differential decode LL3 and dequantize LL3 in a single pass.
// Mathematical identity: cumsum(x) * 2^s == cumsum_of(x * 2^s)
// so we can left-shift each element before accumulating.
if quant.LL3 > 1 {
shift := quant.LL3 - 1
coeffs[4032] <<= shift
for i := 4033; i < 4096; i++ {
coeffs[i] = coeffs[i-1] + coeffs[i]<<shift
}
} else {
for i := 4033; i < 4096; i++ {
coeffs[i] += coeffs[i-1]
}
}
// 3. Dequantize all subbands except LL3 (handled above)
rfxDequantizeSkipLL3(coeffs, quant)
// 4. Inverse DWT (3 levels)
rfxInverseDWT2D(coeffs)
return coeffs
}
// rfxDequantizeSkipLL3 applies dequantization per subband, skipping LL3
// (which is handled together with differential decode in rfxDecodeComponent).
func rfxDequantizeSkipLL3(coeffs []int16, q rfxQuant) {
rfxShiftSubband(coeffs[0:1024], q.HL1) // HL1
rfxShiftSubband(coeffs[1024:2048], q.LH1) // LH1
rfxShiftSubband(coeffs[2048:3072], q.HH1) // HH1
rfxShiftSubband(coeffs[3072:3328], q.HL2) // HL2
rfxShiftSubband(coeffs[3328:3584], q.LH2) // LH2
rfxShiftSubband(coeffs[3584:3840], q.HH2) // HH2
rfxShiftSubband(coeffs[3840:3904], q.HL3) // HL3
rfxShiftSubband(coeffs[3904:3968], q.LH3) // LH3
rfxShiftSubband(coeffs[3968:4032], q.HH3) // HH3
}
// rfxInverseDWT2D performs 3-level inverse 2D discrete wavelet transform in-place.
// Buffer layout: [HL1(1024)|LH1(1024)|HH1(1024)|HL2(256)|LH2(256)|HH2(256)|HL3(64)|LH3(64)|HH3(64)|LL3(64)]
// A single temporary buffer is obtained from the pool and reused across all three
// levels, reducing pool pressure from 9 Get/Put calls (3 levels × 3 components) to 3.
func rfxInverseDWT2D(coeffs []int16) {
bufs := idwtBufPool.Get().(*idwtBufs)
// Level 3: 8×8 subbands → 16×16 output (needs 16×16 = 256 elements)
rfxIDWT2DLevel(coeffs[3840:], bufs.tmp[:256], 8)
// Level 2: 16×16 subbands → 32×32 output (needs 32×32 = 1024 elements)
rfxIDWT2DLevel(coeffs[3072:], bufs.tmp[:1024], 16)
// Level 1: 32×32 subbands → 64×64 output (needs 64×64 = 4096 elements)
rfxIDWT2DLevel(coeffs[0:], bufs.tmp[:4096], 32)
idwtBufPool.Put(bufs)
}
// rfxIDWT2DLevel performs one level of inverse 2D DWT.
// buf contains [HL(n²)|LH(n²)|HH(n²)|LL(n²)] and is replaced with the (2n)×(2n) result.
// tmp is a caller-supplied scratch buffer of length (2n)² (must be ≥ 4n² elements).
// Uses the MS-RDPRFX lifting scheme. Order: horizontal IDWT first, then vertical.
func rfxIDWT2DLevel(buf, tmp []int16, n int) {
nn := n * n
size := 2 * n
// Read subbands directly from buf — no copy needed because the horizontal
// pass only reads from them and writes exclusively to tmp.
hl := buf[0:nn]
lh := buf[nn : 2*nn]
hh := buf[2*nn : 3*nn]
ll := buf[3*nn : 4*nn]
// Step 1: Horizontal IDWT on each row (fused even+odd passes).
// Instead of two separate loops — even pass writing to tmp, then odd pass
// re-reading those values — we keep the last even value in a register and
// compute the preceding odd value in the same iteration. This eliminates
// 2*(n-1) reads of tmp per row (one even[col-1] and one even[col] per odd
// position), replacing them with register references.
// Valid sizes in practice: n = 8, 16, 32.
for row := range n {
rowOff := row * n
lDstOff := row * size
hDstOff := (row + n) * size
// col=0: even boundary (no left neighbour, hl[-1] = hl[0]).
prevEvenL := ll[rowOff] - int16((int32(hl[rowOff])*2+1)>>1)
prevEvenH := lh[rowOff] - int16((int32(hh[rowOff])*2+1)>>1)
tmp[lDstOff] = prevEvenL
tmp[hDstOff] = prevEvenH
// col=1..n-1: compute even[col], then immediately compute odd[col-1]
// using prevEven (=even[col-1], still in register) and the just-computed
// even[col] — no re-read of tmp required.
for col := 1; col < n; col++ {
x := col << 1
evenL := ll[rowOff+col] - int16((int32(hl[rowOff+col-1])+int32(hl[rowOff+col])+1)>>1)
evenH := lh[rowOff+col] - int16((int32(hh[rowOff+col-1])+int32(hh[rowOff+col])+1)>>1)
tmp[lDstOff+x-1] = int16((int32(hl[rowOff+col-1])<<1) + ((int32(prevEvenL)+int32(evenL))>>1))
tmp[hDstOff+x-1] = int16((int32(hh[rowOff+col-1])<<1) + ((int32(prevEvenH)+int32(evenH))>>1))
tmp[lDstOff+x] = evenL
tmp[hDstOff+x] = evenH
prevEvenL = evenL
prevEvenH = evenH
}
// last odd[n-1]: right boundary, even[n] = even[n-1].
x := (n - 1) << 1
tmp[lDstOff+x+1] = int16((int32(hl[rowOff+n-1])<<1) + int32(prevEvenL))
tmp[hDstOff+x+1] = int16((int32(hh[rowOff+n-1])<<1) + int32(prevEvenH))
}
// Step 2: Vertical IDWT on each column.
// Process 8 columns at a time to improve cache utilisation — a cache line
// holds 32 int16 values; 8 columns keeps the working set within one or two
// lines per row access. All valid sizes (16, 32, 64) divide evenly by 8,
// so the scalar tail loop is never reached in practice.
const blk = 8
col := 0
for ; col+blk <= size; col += blk {
// Row 0: first even output (no previous odd)
l0 := tmp[col : col+blk]
h0 := tmp[n*size+col : n*size+col+blk]
out0 := buf[col : col+blk]
for b := range blk {
out0[b] = int16(int32(l0[b]) - ((int32(h0[b])*2 + 1) >> 1))
}
// Rows 1..n-1: interleaved even/odd outputs
for row := 1; row < n; row++ {
lBase := row*size + col
hBase := (row+n)*size + col
hPrevBase := (row-1+n)*size + col
evenBase := 2*row*size + col
prevEvenBase := (2*row-2)*size + col
oddBase := (2*row-1)*size + col
l := tmp[lBase : lBase+blk]
h := tmp[hBase : hBase+blk]
hPrev := tmp[hPrevBase : hPrevBase+blk]
evenOut := buf[evenBase : evenBase+blk]
prevEvenIn := buf[prevEvenBase : prevEvenBase+blk]
oddOut := buf[oddBase : oddBase+blk]
for b := range blk {
hPrevV := int32(hPrev[b])
even := int32(l[b]) - ((hPrevV + int32(h[b]) + 1) >> 1)
evenOut[b] = int16(even)
oddOut[b] = int16((hPrevV << 1) + ((int32(prevEvenIn[b]) + even) >> 1))
}
}
// Last odd row
lastEvenBase := (2*n-2)*size + col
lastHBase := (2*n-1)*size + col
lastEvenSlice := buf[lastEvenBase : lastEvenBase+blk]
lastHSlice := tmp[lastHBase : lastHBase+blk]
lastOddOut := buf[lastHBase : lastHBase+blk]
for b := range blk {
lastOddOut[b] = int16((int32(lastHSlice[b]) << 1) + int32(lastEvenSlice[b]))
}
}
for ; col < size; col++ {
lVal := int32(tmp[col])
hVal := int32(tmp[n*size+col])
buf[col] = int16(lVal - ((hVal*2 + 1) >> 1))
for row := 1; row < n; row++ {
lIdx := row*size + col
hIdx := (row+n)*size + col
hPrevIdx := (row-1+n)*size + col
even := int32(tmp[lIdx]) - ((int32(tmp[hPrevIdx]) + int32(tmp[hIdx]) + 1) >> 1)
buf[2*row*size+col] = int16(even)
prevEven := int32(buf[(2*row-2)*size+col])
odd := (int32(tmp[hPrevIdx]) << 1) + ((prevEven + even) >> 1)
buf[(2*row-1)*size+col] = int16(odd)
}
lastEven := int32(buf[(2*n-2)*size+col])
lastH := int32(tmp[(2*n-1)*size+col])
buf[(2*n-1)*size+col] = int16((lastH << 1) + lastEven)
}
}
func rfxShiftSubband(data []int16, factor uint8) {
if factor <= 1 {
return
}
shift := factor - 1
for i := range data {
data[i] <<= shift
}
}