diff --git a/training/inspect_tflite.py b/training/inspect_tflite.py deleted file mode 100644 index e699997..0000000 --- a/training/inspect_tflite.py +++ /dev/null @@ -1,94 +0,0 @@ -#!/usr/bin/env python3 -"""纯标准库解析 TFLite flatbuffer,打印 subgraph[0] 输入/输出张量的形状与类型。 - -schema 依据 tensorflow/lite/schema/schema.fbs(字段顺序敏感)。 -""" -import struct -import sys - -TENSOR_TYPE = {0: "FLOAT32", 1: "FLOAT16", 2: "INT32", 3: "UINT8", 4: "INT64", - 5: "STRING", 6: "BOOL", 7: "INT16", 8: "COMPLEX64", 9: "INT8", - 10: "FLOAT64", 11: "COMPLEX128", 12: "UINT64", 13: "RESOURCE", - 14: "VARIANT", 15: "UINT32", 16: "UINT16", 17: "INT4", - 18: "BFLOAT16", 19: "FLOAT8_E4M3FN", 20: "FLOAT8_E4M3FNUZ", - 21: "FLOAT8_E5M2", 22: "FLOAT8_E5M2FNUZ"} - - -class FB: - def __init__(self, data): - self.d = data - - def u32(self, pos): - return struct.unpack_from(" vtable + vsize: - return None - foff = self.u16(off) - return None if foff == 0 else table_pos + foff - - def deref(self, pos): - """uoffset 字段解引用:目标 = uoffset 所在位置 + 其值""" - return pos + self.u32(pos) if pos is not None else None - - def vec(self, pos): - pos = self.deref(pos) - if pos is None: - return None - n = self.u32(pos) - return pos + 4, n - - def vec_table(self, pos): - """返回 [table_pos...],每个元素是相对自身位置的 uoffset""" - start, n = self.vec(pos) - return [start + i * 4 + self.u32(start + i * 4) for i in range(n)] - - def string(self, pos): - pos = self.deref(pos) - if pos is None: - return None - n = self.u32(pos) - return self.d[pos + 4:pos + 4 + n].decode("utf-8", "replace") - - def int_vec(self, pos): - start, n = self.vec(pos) - return [self.i32(start + 4 * i) for i in range(n)] - - -def main(path): - fb = FB(open(path, "rb").read()) - root = fb.u32(0) # Model 表 - subgraphs = fb.vec_table(fb.field(root, 2)) # Model.subgraphs [2] - sg = subgraphs[0] - tensors = fb.vec_table(fb.field(sg, 0)) # SubGraph.tensors [0] - inputs = fb.int_vec(fb.field(sg, 1)) # SubGraph.inputs [1] - outputs = fb.int_vec(fb.field(sg, 2)) # SubGraph.outputs [2] - - def describe(i): - t = tensors[i] - shape = fb.int_vec(fb.field(t, 0)) - ttype = fb.u8(fb.field(t, 1)) if fb.field(t, 1) else 0 - name = fb.string(fb.field(t, 3)) or "?" - return f"{name} shape={shape} type={TENSOR_TYPE.get(ttype, ttype)}" - - print("inputs:") - for i in inputs: - print(" ", describe(i)) - print("outputs:") - for i in outputs: - print(" ", describe(i)) - - -if __name__ == "__main__": - main(sys.argv[1])