File size: 5,169 Bytes
f978d83 557735b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 | ---
language:
- en
license: apache-2.0
tags:
- image-classification
- tflite
- flutter
- quickdraw
- doodle-recognition
- on-device-inference
- se-resnet
datasets:
- google/quickdraw
metrics:
- accuracy
library_name: tflite
pipeline_tag: image-classification
model-index:
- name: quickdraw-345-se-resnet
results:
- task:
type: image-classification
name: Image Classification
dataset:
name: Google Quick Draw
type: google/quickdraw
metrics:
- type: accuracy
value: 0.7619
name: Top-1 Accuracy
- type: accuracy
value: 0.8951
name: Top-3 Accuracy
- type: accuracy
value: 0.9226
name: Top-5 Accuracy
- type: accuracy
value: 0.9455
name: Top-10 Accuracy
- type: accuracy
value: 0.7640
name: TFLite Float16 Accuracy
---
# QuickDraw 345 Doodle Classifier — TFLite
A doodle recognition model trained on all **345 categories** from Google's [Quick Draw Dataset](https://quickdraw.withgoogle.com/data), exported as TFLite for **Flutter on-device offline inference**.
## Model Performance
| Metric | Accuracy |
|--------|----------|
| Top-1 | **76.19%** |
| Top-3 | **89.51%** |
| Top-5 | **92.26%** |
| Top-10 | **94.55%** |
| TFLite (float16) | **76.40%** |
> State-of-the-art for 345-class Quick Draw classification is ~73-75% top-1. This model exceeds that.
## Files
| File | Size | Description |
|------|------|-------------|
| `quickdraw_model.tflite` | 8.44 MB | Float16 quantized — recommended for Flutter |
| `quickdraw_model_int8.tflite` | 4.34 MB | Int8 quantized — smallest, fastest |
| `labels.txt` | 2.7 KB | 345 class labels, one per line (alphabetically sorted) |
| `model_metadata.json` | 6.5 KB | Full metadata including accuracy, input shape, Flutter usage |
| `training_history.json` | 5.7 KB | Loss/accuracy per epoch |
| `categories.txt` | 2.7 KB | Raw category list |
## Architecture
- **SE-ResNet** (Squeeze-and-Excitation + ResNet blocks)
- 3 stages: 64 → 128 → 256 filters
- Input: 28×28 grayscale images
- Output: 345-class softmax
- ~3M parameters
## Training
- **Dataset**: Google Quick Draw numpy bitmaps (GCS), 8,000 samples/class × 345 classes = 2.76M images
- **Augmentation**: Random rotation ±8%, translation ±8%, zoom -5%/+10%
- **Optimizer**: Adam + Warmup Cosine Decay
- **Training time**: ~10.9 hours on Kaggle GPU P100
## Flutter Integration
### pubspec.yaml
```yaml
dependencies:
tflite_flutter: ^0.10.4
flutter:
assets:
- assets/quickdraw_model.tflite
- assets/labels.txt
```
### Dart Usage
```dart
import 'package:tflite_flutter/tflite_flutter.dart';
class QuickDrawClassifier {
late Interpreter _interpreter;
late List<String> _labels;
Future<void> load() async {
_interpreter = await Interpreter.fromAsset('assets/quickdraw_model.tflite');
final labelsData = await rootBundle.loadString('assets/labels.txt');
_labels = labelsData.trim().split('\n');
}
/// [pixels] must be a 28x28 Float32List, values in [0.0, 1.0]
/// where 0.0 = black stroke, 1.0 = white background
List<MapEntry<String, double>> predict(Float32List pixels, {int topK = 5}) {
// Reshape to [1, 28, 28, 1]
var input = pixels.reshape([1, 28, 28, 1]);
var output = List.filled(1 * 345, 0.0).reshape([1, 345]);
_interpreter.run(input, output);
final probs = List<double>.from(output[0]);
final indexed = probs.asMap().entries.toList()
..sort((a, b) => b.value.compareTo(a.value));
return indexed.take(topK)
.map((e) => MapEntry(_labels[e.key], e.value))
.toList();
}
}
```
### Preprocessing a drawing canvas
```dart
/// Convert your drawing canvas to a 28x28 normalized Float32List
Float32List canvasToInput(ui.Image image) async {
// Resize to 28x28
final recorder = ui.PictureRecorder();
final canvas = Canvas(recorder);
canvas.drawImageRect(
image,
Rect.fromLTWH(0, 0, image.width.toDouble(), image.height.toDouble()),
Rect.fromLTWH(0, 0, 28, 28),
Paint(),
);
final resized = await recorder.endRecording().toImage(28, 28);
final bytes = await resized.toByteData(format: ui.ImageByteFormat.rawRgba);
// Convert RGBA to grayscale float32, normalize to [0,1]
// white background = 1.0, black strokes = 0.0
final pixels = Float32List(28 * 28);
for (int i = 0; i < 28 * 28; i++) {
final r = bytes!.getUint8(i * 4);
final g = bytes.getUint8(i * 4 + 1);
final b = bytes.getUint8(i * 4 + 2);
pixels[i] = (0.299 * r + 0.587 * g + 0.114 * b) / 255.0;
}
return pixels;
}
```
## Input/Output Spec
| Property | Value |
|----------|-------|
| Input shape | `[1, 28, 28, 1]` |
| Input dtype | `float32` |
| Input range | `[0.0, 1.0]` |
| Background | `1.0` (white) |
| Stroke | `0.0` (black) |
| Output shape | `[1, 345]` |
| Output dtype | `float32` |
| Output | Softmax probabilities |
## License
Model weights: Apache 2.0
Dataset: [Creative Commons Attribution 4.0](https://creativecommons.org/licenses/by/4.0/) (Google Quick Draw)
|