Team Ai
Apppublic

wkea/blockdiffusion-api

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
image_io.py40 linesDownload Raw Back to utils
1from __future__ import annotations
2
3import os
4from typing import Optional
5
6import torch
7from PIL import Image
8from torchvision.utils import make_grid
9
10
11def save_tensor_grid(
12    x: torch.Tensor,
13    path: str,
14    nrow: int,
15    padding: int = 1,
16) -> None:
17    """
18    保存图片网格,支持 RGB(3) / RGBA(4)。
19
20    约定:
21    - x: (B,C,H,W) 且值域在 [0,1]
22    - path: 输出文件路径(推荐 .png)
23    """
24    if x.ndim != 4:
25        raise ValueError("x 必须是 (B,C,H,W)")
26    if x.shape[1] not in (3, 4):
27        raise ValueError("仅支持 3(RGB) 或 4(RGBA) 通道")
28
29    os.makedirs(os.path.dirname(path), exist_ok=True)
30
31    grid = make_grid(x, nrow=nrow, padding=padding)  # (C,H,W)
32    grid = grid.clamp(0.0, 1.0)
33    arr = (grid.permute(1, 2, 0) * 255.0).to(torch.uint8).cpu().numpy()
34
35    mode = "RGBA" if arr.shape[2] == 4 else "RGB"
36    im = Image.fromarray(arr, mode=mode)
37    im.save(path)
38
39
40