fix: rewrite Math preprocessing and decode to match original AntiCAP implementation
Build and Deploy / build (push) Successful in 2m42s
Build and Deploy / deploy (push) Successful in 9s

This commit is contained in:
2026-07-17 10:33:04 +00:00
parent 0c2225bdbb
commit 8c8787135d
+95 -34
View File
@@ -6,6 +6,7 @@ import (
"encoding/json"
"fmt"
"image"
"image/color"
"math"
"net/http"
"os"
@@ -323,14 +324,42 @@ func (h *Handler) Math(imageBase64 string) (string, error) {
}
func preprocessMath(img image.Image) ([]float32, error) {
resized := imaging.Resize(img, 200, 70, imaging.Lanczos)
rgb := imaging.Clone(resized)
// resize_with_padding: 保持比例缩放,白色填充
targetW, targetH := 200, 70
// 计算缩放比例
bounds := img.Bounds()
srcW, srcH := bounds.Dx(), bounds.Dy()
ratio := min(float64(targetW)/float64(srcW), float64(targetH)/float64(srcH))
newW := int(float64(srcW) * ratio)
newH := int(float64(srcH) * ratio)
// 缩放图片
resized := imaging.Resize(img, newW, newH, imaging.Lanczos)
// 创建白色背景
canvas := image.NewRGBA(image.Rect(0, 0, targetW, targetH))
white := color.RGBA{255, 255, 255, 255}
for x := 0; x < targetW; x++ {
for y := 0; y < targetH; y++ {
canvas.Set(x, y, white)
}
}
// 粘贴缩放后的图片
for x := 0; x < newW; x++ {
for y := 0; y < newH; y++ {
canvas.Set(x, y, resized.At(x, y))
}
}
pixels := make([]float32, 3*70*200)
for y := 0; y < 70; y++ {
for x := 0; x < 200; x++ {
c := rgb.At(x, y)
c := canvas.At(x, y)
r, g, b, _ := c.RGBA()
// 归一化到 [0, 1],然后标准化到 [-1, 1]
pixels[0*70*200+y*200+x] = (float32(r)/65535.0 - 0.5) / 0.5
pixels[1*70*200+y*200+x] = (float32(g)/65535.0 - 0.5) / 0.5
pixels[2*70*200+y*200+x] = (float32(b)/65535.0 - 0.5) / 0.5
@@ -341,42 +370,67 @@ func preprocessMath(img image.Image) ([]float32, error) {
}
func decodeMath(output []float32, outputShape []int64) string {
// 输出格式: [num_chars, batch, timesteps] = [51, 1, 19]
// 需要转置为 [batch, timesteps, num_chars] = [1, 19, 51]
// 原版逻辑:
// if preds.shape[1] == 1: # Batch size is 1 at dim 1 implies [T, B, C]
// preds = np.transpose(preds, (1, 0, 2)) # [B, T, C]
// preds_idx = np.argmax(preds, axis=2) # [B, T]
if len(outputShape) < 3 || len(output) == 0 {
return ""
}
numChars := int(outputShape[0]) // 51
batch := int(outputShape[1]) // 1
timesteps := int(outputShape[2]) // 19
// 输出可能是 [T, B, C] 或 [C, B, T]
// 判断依据:如果 shape[1] == 1,说明 batch 在中间,格式是 [T, B, C]
// 否则格式是 [C, B, T] 或 [num_chars, batch, timesteps]
// 处理动态维度(可能返回 -1
if numChars <= 0 {
numChars = 51 // 默认字符数
}
if batch <= 0 {
batch = 1
}
if timesteps <= 0 {
// 从输出大小推断
timesteps = len(output) / (numChars * batch)
}
dim0 := int(outputShape[0])
dim1 := int(outputShape[1])
dim2 := int(outputShape[2])
// 边界检查
if numChars <= 0 || batch <= 0 || timesteps <= 0 {
return ""
}
// 处理动态维度
if dim0 <= 0 { dim0 = 51 }
if dim1 <= 0 { dim1 = 1 }
if dim2 <= 0 { dim2 = len(output) / (dim0 * dim1) }
// 转置: output[c*batch*t + b*t + t] -> transposed[b*timesteps*numChars + t*numChars + c]
transposed := make([]float32, batch*timesteps*numChars)
for b := 0; b < batch; b++ {
var timesteps, batch, numChars int
var data []float32
// 判断输出格式
if dim1 == 1 {
// 格式: [T, B, C] = [19, 1, 51]
// 需要转置为 [B, T, C] = [1, 19, 51]
timesteps = dim0
batch = dim1
numChars = dim2
data = make([]float32, batch*timesteps*numChars)
for t := 0; t < timesteps; t++ {
for c := 0; c < numChars; c++ {
srcIdx := c*batch*timesteps + b*timesteps + t
dstIdx := b*timesteps*numChars + t*numChars + c
if srcIdx < len(output) && dstIdx < len(transposed) {
transposed[dstIdx] = output[srcIdx]
for b := 0; b < batch; b++ {
for c := 0; c < numChars; c++ {
srcIdx := t*batch*numChars + b*numChars + c
dstIdx := b*timesteps*numChars + t*numChars + c
if srcIdx < len(output) && dstIdx < len(data) {
data[dstIdx] = output[srcIdx]
}
}
}
}
} else {
// 格式: [C, B, T] = [51, 1, 19]
// 需要转置为 [B, T, C] = [1, 19, 51]
numChars = dim0
batch = dim1
timesteps = dim2
data = make([]float32, batch*timesteps*numChars)
for c := 0; c < numChars; c++ {
for b := 0; b < batch; b++ {
for t := 0; t < timesteps; t++ {
srcIdx := c*batch*timesteps + b*timesteps + t
dstIdx := b*timesteps*numChars + t*numChars + c
if srcIdx < len(output) && dstIdx < len(data) {
data[dstIdx] = output[srcIdx]
}
}
}
}
@@ -384,7 +438,7 @@ func decodeMath(output []float32, outputShape []int64) string {
// CTC 解码: 对每个 timestep 取 argmax
result := ""
lastIdx := 0
lastIdx := -1
for t := 0; t < timesteps; t++ {
maxIdx := 0
@@ -392,8 +446,8 @@ func decodeMath(output []float32, outputShape []int64) string {
for c := 0; c < numChars; c++ {
idx := t*numChars + c
if idx < len(transposed) && transposed[idx] > maxProb {
maxProb = transposed[idx]
if idx < len(data) && data[idx] > maxProb {
maxProb = data[idx]
maxIdx = c
}
}
@@ -410,6 +464,13 @@ func decodeMath(output []float32, outputShape []int64) string {
return result
}
func min(a, b float64) float64 {
if a < b {
return a
}
return b
}
func evalMathExpression(expr string) (interface{}, error) {
expr = strings.ReplaceAll(expr, "×", "*")
expr = strings.ReplaceAll(expr, "÷", "/")