7.4. Memuat Model#
ml.Model membaca file model dari flash, menguraikannya, mengalokasikan RAM yang dibutuhkan jaringan selama inferensi, dan mengembalikan objek yang membawa semua informasi yang dibutuhkan skrip tentang jaringan yang dimuat.
7.4.1. Konstruktor#
Konstruktor mengambil sebuah path dan post-processor opsional:
model = ml.Model("/rom/blazeface_front_128.tflite",
postprocess=BlazeFace())
Model pada /rom/ (filesystem yang berada di flash) dibaca langsung di tempat: bobot jaringan tetap berada di flash dan model yang dimuat hanya menggunakan RAM sebesar tensor arena. Model pada /sdcard/ disalin ke dalam RAM saat dimuat, sehingga total biayanya adalah ukuran file model ditambah tensor arena. Kedua path berfungsi; pertimbangannya adalah RAM.
Jika file .txt saudara dengan basename yang sama ada, isinya dimuat ke dalam labels secara otomatis. Kata kunci postprocess= mendaftarkan callable yang dijalankan oleh predict() setelah setiap inferensi.
7.4.2. Properti Hanya-Baca#
Model yang dimuat mengekspos sekumpulan kecil properti hanya-baca yang mendeskripsikan jaringan tanpa perlu menjalankannya.
File dan memori.
len-- ukuran file model di disk, dalam byte.ram-- ukuran tensor arena yang dibutuhkan jaringan untuk aktivasi intermedietnya selama inferensi, dalam byte.
Tensor input.
input_shape-- daftar tuple, satu per tensor input, yang memberikan bentuk yang diharapkan oleh jaringan. Jaringan visi memiliki satu input dengan bentuk(1, H, W, C).input_dtype-- daftar kode dtype satu karakter ('b'int8,'B'uint8,'h'int16,'H'uint16,'f'float32), satu per input.input_scaledaninput_zero_point-- parameter kuantisasi yang mengkonversi antara input bernilai nyata tempat jaringan dilatih dan representasi integer yang dijalankan oleh kamera.
Tensor output. Cerminan dari set input: output_shape, output_dtype, output_scale, output_zero_point. Jaringan deteksi menghasilkan dua atau tiga tensor output (kotak, skor kepercayaan, kadang probabilitas kelas); jaringan klasifikasi menghasilkan satu.
Ekstra. labels adalah daftar nama kelas yang dimuat dari file .txt saudara, atau None. postprocess adalah post-processor yang terdaftar, atau None.
7.4.3. Memeriksa BlazeFace#
Memuat model BlazeFace yang disertakan dan mencetak setiap properti memberikan angka yang sebenarnya:
import ml
from ml.postprocessing.mediapipe import BlazeFace
model = ml.Model("/rom/blazeface_front_128.tflite",
postprocess=BlazeFace())
print("file size: ", model.len, "bytes")
print("tensor arena: ", model.ram, "bytes")
print("input shape: ", model.input_shape)
print("input dtype: ", model.input_dtype)
print("input scale: ", model.input_scale)
print("input zp: ", model.input_zero_point)
print("output shape: ", model.output_shape)
print("output dtype: ", model.output_dtype)
print("output scale: ", model.output_scale)
print("output zp: ", model.output_zero_point)
Angka-angka tersebut mengidentifikasi antarmuka jaringan secara konkret: satu tensor input int8 berukuran (1, 128, 128, 3) dan dua output int8 -- satu untuk koefisien regresi kotak, satu untuk skor kepercayaan per-anchor. Parameter kuantisasi mendeskripsikan bagaimana nilai int8 tersebut dipetakan ke float nyata yang dilatih jaringan; post-processor menggunakannya untuk membatalkan kuantisasi sebelum mendekode kotak-kotak.
Setiap properti adalah sumber kebenaran tunggal untuk apa yang dideskripsikannya. Skrip membaca input_shape untuk mengetahui resolusi yang harus diambil, membaca output_scale dan output_zero_point untuk mendekode tensor secara manual, dan membaca labels untuk nama kelas yang dapat dibaca manusia -- tidak pernah di-hardcode, tidak pernah diasumsikan.