Views
No views yet
segmentation-models-pytorch)+ ResNet50 编码器[B, 3, H, W],数值范围 [0, 1][B, 3, H, W],数值范围 [0, 1](末端 sigmoid)config.json:模型结构与训练超参数(导出时写入)model.safetensors:推理用权重(推荐)best.ckpt:原始 PyTorch Lightning checkpoint(用于继续训练/复现实验)configuration.json:简要元数据(framework/task)torch、torchvision、segmentation-models-pytorch、safetensors,以及 Pillow(读写图片可选)。pip install torch torchvision segmentation-models-pytorch safetensors pillow1import json
2from pathlib import Path
3
4import torch
5import segmentation_models_pytorch as smp
6from safetensors.torch import load_file
7from PIL import Image
8import torchvision.transforms.functional as TF
9
10device = "cuda" if torch.cuda.is_available() else "cpu"
11
12# 1) 读取配置
13cfg = json.loads(Path("config.json").read_text(encoding="utf-8"))
14
15# 2) 构建网络(与导出配置保持一致)
16model = smp.UnetPlusPlus(
17 encoder_name=cfg["encoder_name"],
18 encoder_weights=None, # 权重来自 model.safetensors
19 in_channels=cfg["in_channels"],
20 classes=cfg["classes"],
21 decoder_attention_type=cfg.get("decoder_attention_type"),
22 activation=cfg.get("activation"), # 通常为 "sigmoid"
23).to(device)
24
25# 3) 加载权重
26# 说明:导出时可能混入非网络权重(例如 `edge_loss.kx/ky`),推理只需要 Unet++ 本体参数,过滤掉即可。
27state_dict = load_file("model.safetensors")
28model_keys = set(model.state_dict().keys())
29state_dict = {k: v for k, v in state_dict.items() if k in model_keys}
30model.load_state_dict(state_dict, strict=True)
31model.eval()
32
33# 4) 准备输入(训练时仅做 0~1 归一化;如需更贴近训练分布可 resize 到 512x512)
34img = Image.open("input.png").convert("RGB")
35x = TF.to_tensor(img).unsqueeze(0).to(device) # [1,3,H,W] in [0,1]
36
37with torch.no_grad():
38 y = model(x).clamp(0, 1) # [1,3,H,W]
39
40out = TF.to_pil_image(y.squeeze(0).cpu())
41out.save("output.png")pad/resize 到合适尺寸(例如 512x512)。python infer_hd.py --model-dir assets/InkErase --input input.png --output output.pngbest.ckpt(继续训练/复现实验)best.ckpt 是 PyTorch Lightning checkpoint,通常需要配合本项目的 InkEraserModel 代码使用,并提供 ResNet50 预训练权重文件(例如 pretrained_weights/resnet50-0676ba61.pth)。1import torch
2from model import InkEraserModel
3
4model = InkEraserModel.load_from_checkpoint(
5 "best.ckpt",
6 weight="pretrained_weights/resnet50-0676ba61.pth",
7)
8model.eval()
9
10with torch.no_grad():
11 y = model(x)config.json)1{
2 "lr": 0.0001,
3 "weight_decay": 0.01,
4 "loss_w_charb": 0.78,
5 "loss_w_ssim": 0.16,
6 "loss_w_edge": 0.06,
7 "use_mask_loss": true,
8 "loss_mask_weight": 10.0,
9 "charbonnier_eps": 0.001
10}