#!/usr/bin/env python3 # -*- coding: utf-8 -*- """滤光片外观质检 · 推理脚本(单文件,只依赖 onnxruntime + numpy + pillow) 用法 ----- 单张图: python predict.py -i /path/to/filter_0001.png 整个目录(递归),输出 xlsx: python predict.py -i /path/to/images/ -o result.xlsx 同时输出缺陷类型: python predict.py -i /path/to/images/ -o result.xlsx --with-type 自定义判废阈值(默认用标定好的零误报阈值): python predict.py -i imgs/ -o r.xlsx --threshold 0.5 输出 CSV 而不是 xlsx(不需要 openpyxl): python predict.py -i imgs/ -o result.csv """ import os, sys, json, argparse, time, csv try: import numpy as np except ImportError: sys.exit('缺少 numpy:pip install numpy') try: import onnxruntime as ort except ImportError: sys.exit('缺少 onnxruntime:pip install onnxruntime') try: from PIL import Image except ImportError: sys.exit('缺少 Pillow:pip install pillow') if getattr(sys, 'frozen', False): HERE = os.path.dirname(sys.executable) # PyInstaller: exe 所在目录 BUNDLED = getattr(sys, '_MEIPASS', HERE) # 打进二进制里的资源 else: HERE = BUNDLED = os.path.dirname(os.path.abspath(__file__)) IMG_EXT = {'.png', '.jpg', '.jpeg', '.bmp', '.tif', '.tiff'} def find_model(name): """先找 exe 旁边(方便换模型),再找打进包里的。""" for d in (HERE, BUNDLED): p = os.path.join(d, name) if os.path.exists(p): return p return os.path.join(BUNDLED, name) # ---------------------------------------------------------------- 预处理 def preprocess(path, size): """等比缩放到 size 内 + 居中补零 + 逐图灰世界白平衡 + 逐图 z-score。 必须与训练时逐字节一致,否则结果不可信。""" im = Image.open(path).convert('RGB') w, h = im.size s = size / max(w, h) nw, nh = max(1, round(w * s)), max(1, round(h * s)) im = im.resize((nw, nh), Image.BILINEAR) canvas = np.zeros((size, size, 3), np.uint8) y0, x0 = (size - nh) // 2, (size - nw) // 2 canvas[y0:y0 + nh, x0:x0 + nw] = np.asarray(im) f = canvas.astype(np.float32) m = f.reshape(-1, 3).mean(0) # 灰世界 f = f * (m.mean() / np.maximum(m, 1e-6)) f = (f - f.mean()) / max(f.std(), 1e-6) # z-score return np.transpose(f, (2, 0, 1)) # HWC -> CHW def softmax(x, axis=-1): e = np.exp(x - x.max(axis=axis, keepdims=True)) return e / e.sum(axis=axis, keepdims=True) # ---------------------------------------------------------------- 模型 class Model: def __init__(self, onnx_path): meta_p = onnx_path.replace('.onnx', '.json') self.meta = json.load(open(meta_p, encoding='utf-8')) if os.path.exists(meta_p) else {} self.size = int(self.meta.get('size', 320)) self.classes = self.meta.get('classes', ['良品', '不良']) so = ort.SessionOptions() so.intra_op_num_threads = min(8, os.cpu_count() or 4) so.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL prov = ['CUDAExecutionProvider', 'CPUExecutionProvider'] \ if 'CUDAExecutionProvider' in ort.get_available_providers() else ['CPUExecutionProvider'] self.sess = ort.InferenceSession(onnx_path, so, providers=prov) self.provider = self.sess.get_providers()[0] self.iname = self.sess.get_inputs()[0].name def run(self, paths, bs=16): out = [] for i in range(0, len(paths), bs): chunk = paths[i:i + bs] x = np.stack([preprocess(p, self.size) for p in chunk]).astype(np.float32) out.append(softmax(self.sess.run(None, {self.iname: x})[0], axis=1)) return np.concatenate(out) if out else np.zeros((0, len(self.classes))) def collect(inp): if os.path.isfile(inp): return [inp] fs = [] for r, _, names in os.walk(inp): for n in sorted(names): if os.path.splitext(n)[1].lower() in IMG_EXT: fs.append(os.path.join(r, n)) return sorted(fs) def write_table(rows, header, out): ext = os.path.splitext(out)[1].lower() if ext in ('.xlsx', '.xls'): try: from openpyxl import Workbook except ImportError: out = os.path.splitext(out)[0] + '.csv' print(' 未装 openpyxl,改写 CSV:' + out) ext = '.csv' else: wb = Workbook(); ws = wb.active; ws.title = '检测结果' ws.append(header) for r in rows: ws.append(r) for i, w in enumerate([38, 12, 12, 12, 12, 60], 1): if i <= len(header): ws.column_dimensions[chr(64 + i)].width = w wb.save(out); return out with open(out, 'w', newline='', encoding='utf-8-sig') as f: w = csv.writer(f); w.writerow(header); w.writerows(rows) return out def main(): ap = argparse.ArgumentParser(description='滤光片外观质检推理') ap.add_argument('-i', '--input', required=True, help='图片文件或目录') ap.add_argument('-o', '--output', default=None, help='输出 xlsx/csv(目录模式下必填)') ap.add_argument('-m', '--model', default=None) ap.add_argument('--with-type', action='store_true', help='同时用四分类模型给出缺陷类型') ap.add_argument('--type-model', default=None) ap.add_argument('--threshold', type=float, default=None, help='判废阈值(不良概率)。默认读 filter_binary.json 里标定好的值') ap.add_argument('--batch', type=int, default=16) a = ap.parse_args() if a.model is None: a.model = find_model('filter_binary.onnx') if a.type_model is None: a.type_model = find_model('filter_4class.onnx') files = collect(a.input) if not files: sys.exit(f'没找到图片:{a.input}') mdl = Model(a.model) thr = a.threshold if a.threshold is not None else float(mdl.meta.get('threshold', 0.5)) print(f'模型 {os.path.basename(a.model)} | {mdl.provider} | 输入 {mdl.size}px | 判废阈值 {thr:.4f}') print(f'待检 {len(files)} 张') t0 = time.time() P = mdl.run(files, a.batch) score = P[:, 1:].sum(1) # 不良概率 verdict = np.where(score >= thr, '不良', '良品') types = None if a.with_type and os.path.exists(a.type_model): tm = Model(a.type_model) T = tm.run(files, a.batch) names = tm.classes types = [names[i] for i in T.argmax(1)] tconf = T.max(1) dt = time.time() - t0 print(f'完成 {len(files)} 张,用时 {dt:.1f}s({len(files)/max(dt,1e-9):.1f} 张/秒)') if len(files) == 1 and not a.output: print(f'\n 文件 : {os.path.basename(files[0])}') print(f' 判定 : {verdict[0]}') print(f' 不良概率: {score[0]:.4f}') if types: print(f' 缺陷类型: {types[0]}(置信度 {tconf[0]:.3f})') return out = a.output or 'result.xlsx' header = ['图片名', '判定', '不良概率'] rows = [] for i, p in enumerate(files): r = [os.path.basename(p), verdict[i], round(float(score[i]), 4)] if types: r += [types[i] if verdict[i] == '不良' else '', round(float(tconf[i]), 3) if verdict[i] == '不良' else ''] r.append(os.path.relpath(p, a.input) if os.path.isdir(a.input) else p) rows.append(r) if types: header += ['缺陷类型', '类型置信度'] header += ['相对路径'] out = write_table(rows, header, out) n_bad = int((verdict == '不良').sum()) print(f'\n 良品 {len(files)-n_bad} 张 | 不良 {n_bad} 张({n_bad/len(files):.1%})') print(f' 结果 -> {out}') if __name__ == '__main__': main()