Montey/php-edge
0
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 