File size: 3,790 Bytes
4fcca49
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c61e4fd
4fcca49
 
 
 
 
 
 
c61e4fd
4fcca49
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
from __future__ import annotations

from pathlib import Path
import numpy as np
import torch


class OrtDitModule(torch.nn.Module):
    """Drop-in replacement for NitroGen's repeated DiT denoise step.

    CUDA tensors are bound directly to ONNX Runtime buffers, avoiding
    GPU->CPU->GPU copies. TensorRT EP is used when available; otherwise CUDA EP.
    """

    def __init__(self, onnx_path: str | Path, *, prefer_tensorrt: bool = True, cache_dir: str | Path = ".ort-cache"):
        super().__init__()
        import onnxruntime as ort

        try:
            ort.preload_dlls(directory="")
        except Exception:
            try:
                ort.preload_dlls()
            except Exception:
                pass

        available = set(ort.get_available_providers())
        cache_dir = Path(cache_dir)
        cache_dir.mkdir(parents=True, exist_ok=True)
        providers = []
        if prefer_tensorrt and "TensorrtExecutionProvider" in available:
            providers.append(("TensorrtExecutionProvider", {
                "device_id": 0,
                "trt_fp16_enable": True,
                "trt_engine_cache_enable": True,
                "trt_engine_cache_path": str(cache_dir / "trt"),
                "trt_timing_cache_enable": True,
                "trt_timing_cache_path": str(cache_dir / "trt-timing"),
                "trt_builder_optimization_level": 5,
                "trt_max_workspace_size": 1073741824,
            }))
        if "CUDAExecutionProvider" in available:
            providers.append(("CUDAExecutionProvider", {
                "device_id": 0,
                "arena_extend_strategy": "kSameAsRequested",
                "do_copy_in_default_stream": True,
                "cudnn_conv_use_max_workspace": True,
                "gpu_mem_limit": 1610612736,
            }))
        providers.append("CPUExecutionProvider")

        opts = ort.SessionOptions()
        opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
        opts.enable_mem_pattern = True
        opts.enable_cpu_mem_arena = True
        self.session = ort.InferenceSession(str(onnx_path), sess_options=opts, providers=providers)
        self.provider = self.session.get_providers()[0]
        self.output_name = self.session.get_outputs()[0].name
        self._io = self.session.io_binding()

    def extra_repr(self) -> str:
        return f"provider={self.provider}"

    def forward(self, hidden_states, encoder_hidden_states, timestep, **_kwargs):
        h = hidden_states.contiguous().to(dtype=torch.float16)
        e = encoder_hidden_states.contiguous().to(dtype=torch.float16)
        t = timestep.contiguous().to(dtype=torch.int64)

        if h.is_cuda and self.provider != "CPUExecutionProvider":
            out = torch.empty(h.shape, device=h.device, dtype=torch.float16)
            io = self._io
            io.clear_binding_inputs()
            io.clear_binding_outputs()
            io.bind_input("hidden_states", "cuda", h.device.index or 0, np.float16, tuple(h.shape), h.data_ptr())
            io.bind_input("encoder_hidden_states", "cuda", e.device.index or 0, np.float16, tuple(e.shape), e.data_ptr())
            io.bind_input("timestep", "cuda", t.device.index or 0, np.int64, tuple(t.shape), t.data_ptr())
            io.bind_output(self.output_name, "cuda", out.device.index or 0, np.float16, tuple(out.shape), out.data_ptr())
            self.session.run_with_iobinding(io)
            return out

        pred = self.session.run([self.output_name], {
            "hidden_states": h.detach().cpu().numpy(),
            "encoder_hidden_states": e.detach().cpu().numpy(),
            "timestep": t.detach().cpu().numpy(),
        })[0]
        return torch.from_numpy(pred).to(hidden_states.device, dtype=torch.float16)