File size: 1,402 Bytes
a2b9fcc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
# DRUNet (KAIR, MIT) nach ONNX. Die Gewichte sind cszns Original; Architektur aus demselben Repo,
# damit die Schlüssel ohne Umbenennen passen.
import sys, torch
sys.path.insert(0, "/private/tmp/claude-501/-Users-ckret-wkspaces-pxlmonk/35a64bed-d9bb-4a13-9d30-2c08f92e2fd2/scratchpad/kair")
from models.network_unet import UNetRes

SP = "/private/tmp/claude-501/-Users-ckret-wkspaces-pxlmonk/35a64bed-d9bb-4a13-9d30-2c08f92e2fd2/scratchpad"
net = UNetRes(in_nc=4, out_nc=3, nc=[64,128,256,512], nb=4, act_mode='R',
              downsample_mode="strideconv", upsample_mode="convtranspose", bias=False)
sd = torch.load(f"{SP}/kair/drunet_color.pth", map_location="cpu", weights_only=True)
net.load_state_dict(sd, strict=True)   # strict: jede Abweichung ist ein Fehler, kein Zufallsgewicht
missing, unexpected = [], []
print("fehlend:", len(missing), " unerwartet:", len(unexpected))
if missing or unexpected:
    print("  ", missing[:3], unexpected[:3])
net.eval()

# 4. Kanal ist die Rauschpegelkarte: konstant sigma/255 über das ganze Bild.
dummy = torch.randn(1, 4, 256, 256)
out = f"{SP}/models/drunet_color.onnx"
torch.onnx.export(
    net, dummy, out,
    input_names=["input"], output_names=["output"],
    dynamic_axes={"input": {0: "b", 2: "h", 3: "w"}, "output": {0: "b", 2: "h", 3: "w"}},
    opset_version=17,
)
import os
print(f"geschrieben: {out}  ({os.path.getsize(out)/1e6:.1f} MB)")