76 lines
2.6 KiB
Dart
76 lines
2.6 KiB
Dart
|
|
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), '');
|
||
|
|
});
|
||
|
|
});
|
||
|
|
}
|