Files
BadNote/test/ctc_decoder_test.dart

76 lines
2.6 KiB
Dart
Raw Normal View History

import 'package:badnote/services/ocr/ctc_decoder.dart';
import 'package:flutter_test/flutter_test.dart';
/// Build a one-hot-ish logits row of [numClasses] with the max at [maxIndex].
List<double> _row(int numClasses, int maxIndex) {
return List<double>.generate(numClasses, (i) => i == maxIndex ? 1.0 : 0.0);
}
void main() {
group('CtcDecoder.decode', () {
test('collapses consecutive repeats and drops blanks (+1 shift)', () {
// charset indices: 1->'a', 2->'b', 3->'c' (blank at 0, shifted by one).
final decoder = CtcDecoder(['a', 'b', 'c'], blankIndex: 0);
const numClasses = 4; // blank + 3 chars
final logits = <List<double>>[
_row(numClasses, 1), // a
_row(numClasses, 1), // a (collapsed)
_row(numClasses, 0), // blank
_row(numClasses, 2), // b
_row(numClasses, 2), // b (collapsed)
_row(numClasses, 3), // c
];
expect(decoder.decode(logits), 'abc');
});
test('empty input yields empty string', () {
final decoder = CtcDecoder(['a', 'b', 'c'], blankIndex: 0);
expect(decoder.decode(<List<double>>[]), '');
});
test('all-blank input yields empty string', () {
final decoder = CtcDecoder(['a', 'b', 'c'], blankIndex: 0);
const numClasses = 4;
final logits = <List<double>>[
_row(numClasses, 0),
_row(numClasses, 0),
_row(numClasses, 0),
];
expect(decoder.decode(logits), '');
});
test('out-of-range indices are skipped', () {
// charset has 2 entries -> valid class indices are 1 and 2. Class index 3
// maps to charset[2] which is out of range and must be skipped.
final decoder = CtcDecoder(['a', 'b'], blankIndex: 0);
const numClasses = 4;
final logits = <List<double>>[
_row(numClasses, 1), // a
_row(numClasses, 3), // out of range -> skipped
_row(numClasses, 2), // b
];
expect(decoder.decode(logits), 'ab');
});
});
group('CtcDecoder.decodeFlat', () {
test('reshapes a flat row-major list and decodes it', () {
final decoder = CtcDecoder(['a', 'b', 'c'], blankIndex: 0);
const numClasses = 4;
const timeSteps = 3;
final flat = <double>[
..._row(numClasses, 1), // a
..._row(numClasses, 0), // blank
..._row(numClasses, 2), // b
];
expect(decoder.decodeFlat(flat, timeSteps, numClasses), 'ab');
});
test('returns empty for non-positive dimensions', () {
final decoder = CtcDecoder(['a'], blankIndex: 0);
expect(decoder.decodeFlat(<double>[1, 0], 0, 2), '');
expect(decoder.decodeFlat(<double>[1, 0], 2, 0), '');
});
});
}