Download inference.py from AndresDev/captCHAD: direct link, hf CLI and curl.
- Browser
- Download file 8.82 kB
-
https://huggingface.co/AndresDev/captCHAD/resolve/main/inference.py
- Command line
-
hf download hf://AndresDev/captCHAD/inference.py
-
curl -L -o inference.py https://huggingface.co/AndresDev/captCHAD/resolve/main/inference.py
8.82 kB
| """ | |
| 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) # (3, 64, 192) | |
| # Normalize to [-1.0, 1.0] | |
| arr = (arr - 127.5) / 127.5 | |
| return arr[np.newaxis, :, :, :] # (1, 3, 64, 192) | |
| 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__)) | |
| # Resolve weights path if not given | |
| 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: # pytorch | |
| 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: | |
| # Dequantize packed int4 weights | |
| 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) # (48, 1, 63) | |
| 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() | |