| """ |
| captCHAD Inference Engine |
| Supports PyTorch, Safetensors, and ONNX Runtime across FP32, FP16, INT8, FP8, and INT4 quantizations. |
| |
| Usage: |
| python inference.py sample.png |
| python inference.py sample.png --engine onnx --quant int8 |
| python inference.py sample.png --engine onnx --quant fp16 |
| python inference.py sample.png --engine safetensors --quant fp8 |
| python inference.py sample.png --engine safetensors --quant int4 |
| """ |
| import os |
| import sys |
| import argparse |
| import time |
| import numpy as np |
| from PIL import Image |
|
|
| CHARSET = "0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ" |
| IDX2CHAR = {i + 1: ch for i, ch in enumerate(CHARSET)} |
| BLANK_IDX = 0 |
|
|
| def preprocess_image(image_path: str) -> np.ndarray: |
| """Preprocess image to normalized float32 array (1, 3, 64, 192).""" |
| with Image.open(image_path) as img: |
| if img.mode == "RGBA": |
| bg = Image.new("RGB", img.size, (255, 255, 255)) |
| bg.paste(img, mask=img.split()[3]) |
| img = bg |
| elif img.mode != "RGB": |
| img = img.convert("RGB") |
| img = img.resize((192, 64), Image.BILINEAR) |
| arr = np.array(img, dtype=np.float32).transpose(2, 0, 1) |
| |
| arr = (arr - 127.5) / 127.5 |
| return arr[np.newaxis, :, :, :] |
|
|
| def ctc_decode_greedy(tokens: list[int]) -> str: |
| """Standard CTC collapse: drop consecutive duplicates and blank tokens.""" |
| res = [] |
| prev = None |
| for t in tokens: |
| if t != prev and t != BLANK_IDX: |
| if t in IDX2CHAR: |
| res.append(IDX2CHAR[t]) |
| prev = t |
| return "".join(res) |
|
|
| class captCHADPredictor: |
| def __init__(self, engine: str = "onnx", quant: str = "fp32", weights_path: str = None): |
| self.engine = engine.lower() |
| self.quant = quant.lower() |
| dir_path = os.path.dirname(os.path.abspath(__file__)) |
|
|
| |
| if weights_path is None: |
| if self.engine == "onnx": |
| if self.quant == "int8": |
| weights_path = os.path.join(dir_path, "captchad_int8.onnx") |
| elif self.quant == "fp16": |
| weights_path = os.path.join(dir_path, "captchad_fp16.onnx") |
| else: |
| weights_path = os.path.join(dir_path, "captchad.onnx") |
| elif self.engine == "safetensors": |
| if self.quant == "fp16": |
| weights_path = os.path.join(dir_path, "model_fp16.safetensors") |
| elif self.quant == "fp8": |
| weights_path = os.path.join(dir_path, "model_fp8.safetensors") |
| elif self.quant == "int4": |
| weights_path = os.path.join(dir_path, "model_int4.safetensors") |
| else: |
| weights_path = os.path.join(dir_path, "model.safetensors") |
| else: |
| if self.quant == "fp16": |
| weights_path = os.path.join(dir_path, "captchad_fp16.pt") |
| elif self.quant == "int8": |
| weights_path = os.path.join(dir_path, "captchad_int8.pt") |
| elif self.quant == "fp8": |
| weights_path = os.path.join(dir_path, "captchad_fp8.pt") |
| elif self.quant == "int4": |
| weights_path = os.path.join(dir_path, "captchad_int4.pt") |
| else: |
| weights_path = os.path.join(dir_path, "captchad.pt") |
|
|
| self.weights_path = weights_path |
|
|
| if self.engine == "onnx": |
| import onnxruntime as ort |
| opts = ort.SessionOptions() |
| opts.intra_op_num_threads = min(4, os.cpu_count() or 4) |
| self.session = ort.InferenceSession(weights_path, sess_options=opts) |
| self.input_name = self.session.get_inputs()[0].name |
| elif self.engine in ("pytorch", "safetensors"): |
| import torch |
| from model import captCHAD, decode_beam_search_single |
| self.torch = torch |
| self.decode_beam = decode_beam_search_single |
| self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
| self.model = captCHAD(num_classes=len(CHARSET) + 1) |
|
|
| if weights_path.endswith("int8.pt"): |
| ckpt = torch.load(weights_path, map_location=self.device, weights_only=False) |
| self.model = ckpt["model"] if "model" in ckpt else ckpt |
| elif weights_path.endswith(".safetensors"): |
| from safetensors.torch import load_file |
| state = load_file(weights_path) |
| if "int4" in weights_path: |
| |
| base_ckpt = torch.load(os.path.join(dir_path, "captchad.pt"), map_location="cpu") |
| orig_shapes = {k: v.shape for k, v in base_ckpt["model_state_dict"].items()} |
| restored = {} |
| for k, v in state.items(): |
| if k.endswith(".packed_int4"): |
| bname = k[:-len(".packed_int4")] |
| sc = state[f"{bname}.scale"].squeeze() |
| oshape = orig_shapes[bname] |
| low = (v & 0x0F).to(torch.int8) - 7 |
| high = ((v >> 4) & 0x0F).to(torch.int8) - 7 |
| unpacked = torch.empty(len(v) * 2, dtype=torch.int8) |
| unpacked[0::2], unpacked[1::2] = low, high |
| nel = 1 |
| for d in oshape: nel *= d |
| restored[bname] = (unpacked[:nel].to(torch.float32) * sc).reshape(oshape) |
| elif k.endswith(".scale"): |
| continue |
| else: |
| restored[k] = v.float() if v.is_floating_point() else v |
| state = restored |
| else: |
| state = {k: v.to(torch.float32) if v.is_floating_point() else v for k, v in state.items()} |
| self.model.load_state_dict(state) |
| else: |
| ckpt = torch.load(weights_path, map_location=self.device) |
| state = ckpt["model_state_dict"] if "model_state_dict" in ckpt else ckpt |
| state = {k: v.to(torch.float32) if v.is_floating_point() else v for k, v in state.items()} |
| self.model.load_state_dict(state) |
|
|
| self.model.to(self.device) |
| self.model.eval() |
|
|
| def predict(self, image_path: str, use_beam: bool = False) -> tuple[str, float]: |
| t0 = time.perf_counter() |
| inp = preprocess_image(image_path) |
|
|
| if self.engine == "onnx": |
| logits = self.session.run(None, {self.input_name: inp})[0] |
| preds = np.argmax(logits[:, 0, :], axis=-1).tolist() |
| text = ctc_decode_greedy(preds) |
| else: |
| t = self.torch.from_numpy(inp).to(self.device) |
| with self.torch.no_grad(): |
| logits = self.model(t) |
| if use_beam: |
| log_probs = logits[:, 0, :].log_softmax(dim=-1) |
| beam_res = self.decode_beam(log_probs, beam_width=15) |
| text = beam_res[0][0] if beam_res else "" |
| else: |
| preds = logits.argmax(dim=-1)[:, 0].tolist() |
| text = ctc_decode_greedy(preds) |
|
|
| latency_ms = (time.perf_counter() - t0) * 1000 |
| return text, latency_ms |
|
|
| def main(): |
| parser = argparse.ArgumentParser(description="captCHAD Multi-Format Inference Engine") |
| parser.add_argument("image", nargs="?", default="sample.png", help="Path to input image") |
| parser.add_argument("--engine", choices=["onnx", "pytorch", "safetensors"], default="onnx", help="Inference engine") |
| parser.add_argument("--quant", choices=["fp32", "fp16", "int8", "fp8", "int4"], default="fp32", help="Precision format") |
| parser.add_argument("--weights", type=str, default=None, help="Custom weights file path") |
| parser.add_argument("--beam", action="store_true", help="Use CTC beam search (PyTorch only)") |
| args = parser.parse_args() |
|
|
| if not os.path.exists(args.image): |
| print(f"Error: image not found at '{args.image}'.") |
| sys.exit(1) |
|
|
| predictor = captCHADPredictor(engine=args.engine, quant=args.quant, weights_path=args.weights) |
| pred_text, latency = predictor.predict(args.image, use_beam=args.beam) |
|
|
| print(f"Image: {args.image}") |
| print(f"Engine: {args.engine.upper()} ({args.quant.upper()})") |
| print(f"Weights: {os.path.basename(predictor.weights_path)}") |
| print(f"Prediction: {pred_text}") |
| print(f"Latency: {latency:.2f} ms") |
|
|
| if __name__ == "__main__": |
| main() |
|
|