| |
| |
| 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) |
| missing, unexpected = [], [] |
| print("fehlend:", len(missing), " unerwartet:", len(unexpected)) |
| if missing or unexpected: |
| print(" ", missing[:3], unexpected[:3]) |
| net.eval() |
|
|
| |
| 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)") |
|
|