Team Ai
Modelpublic

Montey/php-edge

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
deploy.py145 linesDownload Raw Back to trt_pipeline
1# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.2# SPDX-License-Identifier: MIT3#4# Permission is hereby granted, free of charge, to any person obtaining a5# copy of this software and associated documentation files (the "Software"),6# to deal in the Software without restriction, including without limitation7# the rights to use, copy, modify, merge, publish, distribute, sublicense,8# and/or sell copies of the Software, and to permit persons to whom the9# Software is furnished to do so, subject to the following conditions:10#11# The above copyright notice and this permission notice shall be included in12# all copies or substantial portions of the Software.13#14# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR15# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,16# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL17# THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER18# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING19# FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER20# DEALINGS IN THE SOFTWARE.21 22import types23from pathlib import Path24 25import tensorrt as trt26import torch27from cache_diffusion.cachify import CACHED_PIPE, get_model28from cuda import cudart29from diffusers.models.transformers.transformer_sd3 import SD3Transformer2DModel30from diffusers.models.unets.unet_2d_condition import UNet2DConditionModel31from trt_pipeline.config import ONNX_CONFIG32from trt_pipeline.models.sd3 import sd3_forward33from trt_pipeline.models.sdxl import (34    cachecrossattnupblock2d_forward,35    cacheunet_forward,36    cacheupblock2d_forward,37)38from polygraphy.backend.trt import (39    CreateConfig,40    Profile,41    engine_from_network,42    network_from_onnx_path,43    save_engine,44)45from torch.onnx import export as onnx_export46 47from .utils import Engine48 49 50def replace_new_forward(backbone):51    if backbone.__class__ == UNet2DConditionModel:52        backbone.forward = types.MethodType(cacheunet_forward, backbone)53        for upsample_block in backbone.up_blocks:54            if (55                hasattr(upsample_block, "has_cross_attention")56                and upsample_block.has_cross_attention57            ):58                upsample_block.forward = types.MethodType(59                    cachecrossattnupblock2d_forward, upsample_block60                )61            else:62                upsample_block.forward = types.MethodType(cacheupblock2d_forward, upsample_block)63    elif backbone.__class__ == SD3Transformer2DModel:64        backbone.forward = types.MethodType(sd3_forward, backbone)65 66 67def get_input_info(dummy_dict, info: str = None, batch_size: int = 1):68    return_val = [] if info == "profile_shapes" or info == "input_names" else {}69 70    def collect_leaf_keys(d):71        for key, value in d.items():72            if isinstance(value, dict):73                collect_leaf_keys(value)74            else:75                value = (value[0] * batch_size,) + value[1:]76                if info == "profile_shapes":77                    return_val.append((key, value))  # type: ignore78                elif info == "profile_shapes_dict":79                    return_val[key] = value  # type: ignore80                elif info == "dummy_input":81                    return_val[key] = torch.ones(value).half().cuda()  # type: ignore82                elif info == "input_names":83                    return_val.append(key)  # type: ignore84 85    collect_leaf_keys(dummy_dict)86    return return_val87 88 89def get_total_device_memory(backbone):90    max_device_memory = 091    for _, engine in backbone.engines.items():92        max_device_memory = max(max_device_memory, engine.engine.device_memory_size)93    return max_device_memory94 95 96def load_engines(backbone, engine_path: Path, batch_size: int = 1):97    backbone.engines = {}98    for f in engine_path.iterdir():99        if f.is_file():100            eng = Engine()101            eng.load(str(f))102            backbone.engines[f"{f.stem}"] = eng103    _, shared_device_memory = cudart.cudaMalloc(get_total_device_memory(backbone))104    for engine in backbone.engines.values():105        engine.activate(shared_device_memory)106    backbone.cuda_stream = cudart.cudaStreamCreate()[1]107    for block_name in backbone.engines.keys():108        backbone.engines[block_name].allocate_buffers(109            shape_dict=get_input_info(110                ONNX_CONFIG[backbone.__class__][block_name]["dummy_input"],111                "profile_shapes_dict",112                batch_size,113            ),114            device=backbone.device,115            batch_size=batch_size,116        )117    # TODO: Free and clean up the origin pytorch cuda memory118 119 120def warm_up(backbone, batch_size: int = 1):121    print("Warming-up TensorRT engines...")122    for name, engine in backbone.engines.items():123        dummy_input = get_input_info(124            ONNX_CONFIG[backbone.__class__][name]["dummy_input"], "dummy_input", batch_size125        )126        _ = engine(dummy_input, backbone.cuda_stream)127 128 129def teardown(pipe):130    backbone = get_model(pipe)131    for engine in backbone.engines.values():132        del engine133 134    cudart.cudaStreamDestroy(backbone.cuda_stream)135    del backbone.cuda_stream136 137 138def load_unet_trt(unet, engine_path: Path, batch_size: int = 1):139    backbone = unet 140    engine_path.mkdir(parents=True, exist_ok=True)141    replace_new_forward(backbone) 142    load_engines(backbone, engine_path, batch_size)143    warm_up(backbone, batch_size)144    backbone.use_trt_infer = True145