#!/usr/bin/env python """导出 torch.jit(fp16) 512x512 模型: 输入 LR[1,3,512,512][-1,1] -> 输出同尺寸。 内部: bicubic 512->128 -> 官方 4x 学生全链 -> 512。 (AdaIN 后处理不进测速模型) 用法: python src/export_jit.py --net weight/s2/net_params_X.pkl --out model_dir """ import argparse, os, sys from pathlib import Path REPO = Path(__file__).resolve().parents[1] sys.path.insert(0, str(REPO)); sys.path.insert(0, str(REPO / "src")); sys.path.insert(0, str(REPO / "official")) import torch, torch.nn as nn, torch.nn.functional as F from common import load_diffusers_sd, load_pruned_decoder, build_net class SR512(nn.Module): def __init__(self, net, tail): super().__init__() self.net = net self.tail = tail def forward(self, x512): x128 = F.interpolate(x512, size=(128, 128), mode="bicubic", align_corners=False) z = self.net(x128) return self.tail(z) def main(): ap = argparse.ArgumentParser() ap.add_argument("--net", required=True) ap.add_argument("--out", default="model_dir") ap.add_argument("--half_decoder", default="weight/pretrained/halfDecoder.ckpt") ap.add_argument("--model_id", default="models/stable-diffusion-2-1-base") ap.add_argument("--name", default="your_model.pt") args = ap.parse_args() os.makedirs(args.out, exist_ok=True) device = "cuda" if torch.cuda.is_available() else "cpu" vae, unet, _, _ = load_diffusers_sd(args.model_id, dtype=torch.float32, device="cpu") del vae decoder = load_pruned_decoder(args.half_decoder, device="cpu", dtype=torch.float32) net = build_net(unet, decoder) sd = torch.load(args.net, map_location="cpu", weights_only=False) if any(k.startswith("module.") for k in sd): sd = {k.replace("module.", "", 1): v for k, v in sd.items()} net.load_state_dict(sd, strict=True) net.eval() tail = nn.Sequential(*decoder.up_blocks, decoder.conv_norm_out, decoder.conv_act, decoder.conv_out).eval() model = SR512(net, tail).to(device).half().eval() # trace on fixed 512 fp16 dummy = torch.randn(1, 3, 512, 512, device=device).half() * 0.5 with torch.no_grad(): traced = torch.jit.trace(model, dummy, check_trace=False) traced = torch.jit.freeze(traced) out_path = os.path.join(args.out, args.name) traced.save(out_path) # 自检: 两次前向一致性 + 形状 with torch.no_grad(): o1 = traced(dummy); o2 = traced(dummy) assert o1.shape == dummy.shape, o1.shape err = (o1 - o2).abs().max().item() print("saved", out_path, "| deterministic max-diff:", err) print("output range sample:", float(o1.min()), float(o1.max())) if __name__ == "__main__": main()