Update for 0.31.0 release

This commit is contained in:
Keval Morabia
2025-06-05 13:24:07 -07:00
parent 3039f76d6a
commit 7af33d29ce
378 changed files with 10134 additions and 5835 deletions
@@ -39,7 +39,7 @@ def parse_args():
args = parser.parse_args()
for key, value in vars(args).items():
if value is not None:
print("Parsed args -- {}: {}".format(key, value))
print(f"Parsed args -- {key}: {value}")
return args
@@ -134,5 +134,5 @@ def infer(pipe):
def prepare(pipe, config_list):
model = get_model(pipe)
assert model.__class__ in CACHED_PIPE.keys(), f"{model.__class__} is not supported!"
assert model.__class__ in CACHED_PIPE, f"{model.__class__} is not supported!"
cachify(model, config_list, CACHED_PIPE[model.__class__])
@@ -58,8 +58,8 @@ def replace_new_forward(backbone):
backbone.forward = types.MethodType(sd3_forward, backbone)
def get_input_info(dummy_dict, info: str = None, batch_size: int = 1):
return_val = [] if info == "profile_shapes" or info == "input_names" else {}
def get_input_info(dummy_dict, info: str | None = None, batch_size: int = 1):
return_val = [] if info in {"profile_shapes", "input_names"} else {}
def collect_leaf_keys(d):
for key, value in d.items():
@@ -68,13 +68,13 @@ def get_input_info(dummy_dict, info: str = None, batch_size: int = 1):
else:
value = (value[0] * batch_size,) + value[1:]
if info == "profile_shapes":
return_val.append((key, value)) # type: ignore
return_val.append((key, value)) # type: ignore[attr-defined]
elif info == "profile_shapes_dict":
return_val[key] = value # type: ignore
return_val[key] = value # type: ignore[index]
elif info == "dummy_input":
return_val[key] = torch.ones(value).half().cuda() # type: ignore
return_val[key] = torch.ones(value).half().cuda() # type: ignore[index]
elif info == "input_names":
return_val.append(key) # type: ignore
return_val.append(key) # type: ignore[attr-defined]
collect_leaf_keys(dummy_dict)
return return_val
@@ -83,12 +83,12 @@ def get_input_info(dummy_dict, info: str = None, batch_size: int = 1):
def complie2trt(cls, onnx_path: Path, engine_path: Path, batch_size: int = 1):
subdirs = [f for f in onnx_path.iterdir() if f.is_dir()]
for subdir in subdirs:
if subdir.name not in ONNX_CONFIG[cls].keys():
if subdir.name not in ONNX_CONFIG[cls]:
continue
model_path = subdir / "model.onnx"
plan_path = engine_path / f"{subdir.name}.plan"
if not plan_path.exists():
print(f"Building {str(model_path)}")
print(f"Building {model_path!s}")
build_profile = Profile()
profile_shapes = get_input_info(
ONNX_CONFIG[cls][subdir.name]["dummy_input"], "profile_shapes", batch_size
@@ -109,12 +109,12 @@ def complie2trt(cls, onnx_path: Path, engine_path: Path, batch_size: int = 1):
)
save_engine(engine, path=plan_path)
else:
print(f"{str(model_path)} already exists!")
print(f"{model_path!s} already exists!")
def get_total_device_memory(backbone):
max_device_memory = 0
for _, engine in backbone.engines.items():
for engine in backbone.engines.values():
max_device_memory = max(max_device_memory, engine.engine.device_memory_size)
return max_device_memory
@@ -130,7 +130,7 @@ def load_engines(backbone, engine_path: Path, batch_size: int = 1):
for engine in backbone.engines.values():
engine.activate(shared_device_memory)
backbone.cuda_stream = cudart.cudaStreamCreate()[1]
for block_name in backbone.engines.keys():
for block_name in backbone.engines:
backbone.engines[block_name].allocate_buffers(
shape_dict=get_input_info(
ONNX_CONFIG[backbone.__class__][block_name]["dummy_input"],
@@ -178,7 +178,7 @@ def export_onnx(backbone, onnx_path: Path):
opset_version=17,
)
else:
print(f"{str(_onnx_file)} alread exists!")
print(f"{_onnx_file!s} alread exists!")
def warm_up(backbone, batch_size: int = 1):
@@ -13,7 +13,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.
from typing import Any, Optional, Union
from typing import Any
import torch
from diffusers.models.modeling_outputs import Transformer2DModelOutput
@@ -31,10 +31,10 @@ def sd3_forward(
encoder_hidden_states: torch.FloatTensor = None,
pooled_projections: torch.FloatTensor = None,
timestep: torch.LongTensor = None,
block_controlnet_hidden_states: list = None,
joint_attention_kwargs: Optional[dict[str, Any]] = None,
block_controlnet_hidden_states: list | None = None,
joint_attention_kwargs: dict | None = None,
return_dict: bool = True,
) -> Union[torch.FloatTensor, Transformer2DModelOutput]:
) -> torch.FloatTensor | Transformer2DModelOutput:
"""
The [`SD3Transformer2DModel`] forward method.
@@ -100,25 +100,24 @@ def sd3_forward(
**ckpt_kwargs,
)
elif hasattr(self, "use_trt_infer") and self.use_trt_infer:
feed_dict = {
"hidden_states": hidden_states,
"encoder_hidden_states": encoder_hidden_states,
"temb": temb,
}
_results = self.engines[f"transformer_blocks.{index_block}"](
feed_dict, self.cuda_stream
)
if index_block != 23:
encoder_hidden_states = _results["encoder_hidden_states_out"]
hidden_states = _results["hidden_states_out"]
else:
if hasattr(self, "use_trt_infer") and self.use_trt_infer:
feed_dict = {
"hidden_states": hidden_states,
"encoder_hidden_states": encoder_hidden_states,
"temb": temb,
}
_results = self.engines[f"transformer_blocks.{index_block}"](
feed_dict, self.cuda_stream
)
if index_block != 23:
encoder_hidden_states = _results["encoder_hidden_states_out"]
hidden_states = _results["hidden_states_out"]
else:
encoder_hidden_states, hidden_states = block(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
temb=temb,
)
encoder_hidden_states, hidden_states = block(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
temb=temb,
)
# controlnet residual
if block_controlnet_hidden_states is not None and block.context_pre_only is False:
@@ -32,7 +32,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.
from typing import Any, Optional, Union
from typing import Any
import torch
from diffusers.models.unets.unet_2d_condition import UNet2DConditionOutput
@@ -44,12 +44,12 @@ def cachecrossattnupblock2d_forward(
res_hidden_states_0: torch.FloatTensor,
res_hidden_states_1: torch.FloatTensor,
res_hidden_states_2: torch.FloatTensor,
temb: Optional[torch.FloatTensor] = None,
encoder_hidden_states: Optional[torch.FloatTensor] = None,
cross_attention_kwargs: Optional[dict[str, Any]] = None,
upsample_size: Optional[int] = None,
attention_mask: Optional[torch.FloatTensor] = None,
encoder_attention_mask: Optional[torch.FloatTensor] = None,
temb: torch.FloatTensor | None = None,
encoder_hidden_states: torch.FloatTensor | None = None,
cross_attention_kwargs: dict[str, Any] | None = None,
upsample_size: int | None = None,
attention_mask: torch.FloatTensor | None = None,
encoder_attention_mask: torch.FloatTensor | None = None,
) -> torch.FloatTensor:
res_hidden_states_tuple = (res_hidden_states_0, res_hidden_states_1, res_hidden_states_2)
for resnet, attn in zip(self.resnets, self.attentions):
@@ -82,8 +82,8 @@ def cacheupblock2d_forward(
res_hidden_states_0: torch.FloatTensor,
res_hidden_states_1: torch.FloatTensor,
res_hidden_states_2: torch.FloatTensor,
temb: Optional[torch.FloatTensor] = None,
upsample_size: Optional[int] = None,
temb: torch.FloatTensor | None = None,
upsample_size: int | None = None,
) -> torch.FloatTensor:
res_hidden_states_tuple = (res_hidden_states_0, res_hidden_states_1, res_hidden_states_2)
for resnet in self.resnets:
@@ -105,19 +105,19 @@ def cacheupblock2d_forward(
def cacheunet_forward(
self,
sample: torch.FloatTensor,
timestep: Union[torch.Tensor, float, int],
timestep: torch.Tensor | float | int,
encoder_hidden_states: torch.Tensor,
class_labels: Optional[torch.Tensor] = None,
timestep_cond: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
cross_attention_kwargs: Optional[dict[str, Any]] = None,
added_cond_kwargs: Optional[dict[str, torch.Tensor]] = None,
down_block_additional_residuals: Optional[tuple[torch.Tensor]] = None,
mid_block_additional_residual: Optional[torch.Tensor] = None,
down_intrablock_additional_residuals: Optional[tuple[torch.Tensor]] = None,
encoder_attention_mask: Optional[torch.Tensor] = None,
class_labels: torch.Tensor | None = None,
timestep_cond: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
cross_attention_kwargs: dict[str, Any] | None = None,
added_cond_kwargs: dict[str, torch.Tensor] | None = None,
down_block_additional_residuals: tuple[torch.Tensor] | None = None,
mid_block_additional_residual: torch.Tensor | None = None,
down_intrablock_additional_residuals: tuple[torch.Tensor] | None = None,
encoder_attention_mask: torch.Tensor | None = None,
return_dict: bool = True,
) -> Union[UNet2DConditionOutput, tuple]:
) -> UNet2DConditionOutput | tuple:
# 1. time
t_emb = self.get_time_embed(sample=sample, timestep=timestep)
emb = self.time_embedding(t_emb, timestep_cond)
@@ -161,7 +161,7 @@ def cacheunet_forward(
sample = down_results["sample"]
res_samples_0 = down_results["res_samples_0"]
res_samples_1 = down_results["res_samples_1"]
if "res_samples_2" in down_results.keys():
if "res_samples_2" in down_results:
res_samples_2 = down_results["res_samples_2"]
else:
# For t2i-adapter CrossAttnDownBlock2D
@@ -176,24 +176,23 @@ def cacheunet_forward(
encoder_attention_mask=encoder_attention_mask,
**additional_residuals,
)
elif hasattr(self, "use_trt_infer") and self.use_trt_infer:
feed_dict = {"hidden_states": sample, "temb": emb}
down_results = self.engines[f"down_blocks.{i}"](feed_dict, self.cuda_stream)
sample = down_results["sample"]
res_samples_0 = down_results["res_samples_0"]
res_samples_1 = down_results["res_samples_1"]
if "res_samples_2" in down_results:
res_samples_2 = down_results["res_samples_2"]
else:
if hasattr(self, "use_trt_infer") and self.use_trt_infer:
feed_dict = {"hidden_states": sample, "temb": emb}
down_results = self.engines[f"down_blocks.{i}"](feed_dict, self.cuda_stream)
sample = down_results["sample"]
res_samples_0 = down_results["res_samples_0"]
res_samples_1 = down_results["res_samples_1"]
if "res_samples_2" in down_results.keys():
res_samples_2 = down_results["res_samples_2"]
else:
sample, res_samples = downsample_block(hidden_states=sample, temb=emb)
sample, res_samples = downsample_block(hidden_states=sample, temb=emb)
if hasattr(self, "use_trt_infer") and self.use_trt_infer:
down_block_res_samples += (
res_samples_0,
res_samples_1,
)
if "res_samples_2" in down_results.keys():
if "res_samples_2" in down_results:
down_block_res_samples += (res_samples_2,)
else:
down_block_res_samples += res_samples
@@ -245,25 +244,24 @@ def cacheunet_forward(
attention_mask=attention_mask,
encoder_attention_mask=encoder_attention_mask,
)
elif hasattr(self, "use_trt_infer") and self.use_trt_infer:
feed_dict = {
"hidden_states": sample,
"res_hidden_states_0": res_samples[0],
"res_hidden_states_1": res_samples[1],
"res_hidden_states_2": res_samples[2],
"temb": emb,
}
up_results = self.engines[f"up_blocks.{i}"](feed_dict, self.cuda_stream)
sample = up_results["sample"]
else:
if hasattr(self, "use_trt_infer") and self.use_trt_infer:
feed_dict = {
"hidden_states": sample,
"res_hidden_states_0": res_samples[0],
"res_hidden_states_1": res_samples[1],
"res_hidden_states_2": res_samples[2],
"temb": emb,
}
up_results = self.engines[f"up_blocks.{i}"](feed_dict, self.cuda_stream)
sample = up_results["sample"]
else:
sample = upsample_block(
hidden_states=sample,
temb=emb,
res_hidden_states_0=res_samples[0],
res_hidden_states_1=res_samples[1],
res_hidden_states_2=res_samples[2],
)
sample = upsample_block(
hidden_states=sample,
temb=emb,
res_hidden_states_0=res_samples[0],
res_hidden_states_1=res_samples[1],
res_hidden_states_2=res_samples[2],
)
# 6. post-process
if self.conv_norm_out:
@@ -58,22 +58,22 @@ class Engine:
def activate(self, reuse_device_memory=None):
if reuse_device_memory:
self.context = self.engine.create_execution_context_without_device_memory() # type: ignore
self.context = self.engine.create_execution_context_without_device_memory() # type: ignore[union-attr]
self.context.device_memory = reuse_device_memory
else:
self.context = self.engine.create_execution_context() # type: ignore
self.context = self.engine.create_execution_context() # type: ignore[union-attr]
def allocate_buffers(self, shape_dict=None, device="cuda", batch_size=1):
for binding in range(self.engine.num_io_tensors): # type: ignore
name = self.engine.get_tensor_name(binding) # type: ignore
for binding in range(self.engine.num_io_tensors): # type: ignore[union-attr]
name = self.engine.get_tensor_name(binding) # type: ignore[union-attr]
if shape_dict and name in shape_dict:
shape = shape_dict[name]
else:
shape = self.engine.get_tensor_shape(name) # type: ignore
shape = self.engine.get_tensor_shape(name) # type: ignore[union-attr]
shape = (batch_size * 2,) + shape[1:]
dtype = trt.nptype(self.engine.get_tensor_dtype(name)) # type: ignore
if self.engine.get_tensor_mode(name) == trt.TensorIOMode.INPUT: # type: ignore
self.context.set_input_shape(name, shape) # type: ignore
dtype = trt.nptype(self.engine.get_tensor_dtype(name)) # type: ignore[union-attr]
if self.engine.get_tensor_mode(name) == trt.TensorIOMode.INPUT: # type: ignore[union-attr]
self.context.set_input_shape(name, shape) # type: ignore[union-attr]
tensor = torch.empty(tuple(shape), dtype=numpy_to_torch_dtype_dict[dtype]).to(
device=device
)
@@ -84,7 +84,7 @@ class Engine:
self.tensors[name].copy_(buf)
for name, tensor in self.tensors.items():
self.context.set_tensor_address(name, tensor.data_ptr()) # type: ignore
self.context.set_tensor_address(name, tensor.data_ptr()) # type: ignore[union-attr]
if use_cuda_graph:
if self.cuda_graph_instance is not None:
@@ -92,7 +92,7 @@ class Engine:
cuassert(cudart.cudaStreamSynchronize(stream))
else:
# do inference before CUDA graph capture
noerror = self.context.execute_async_v3(stream) # type: ignore
noerror = self.context.execute_async_v3(stream) # type: ignore[union-attr]
if not noerror:
raise ValueError("ERROR: inference failed.")
# capture cuda graph
@@ -101,11 +101,11 @@ class Engine:
stream, cudart.cudaStreamCaptureMode.cudaStreamCaptureModeGlobal
)
)
self.context.execute_async_v3(stream) # type: ignore
self.context.execute_async_v3(stream) # type: ignore[union-attr]
self.graph = cuassert(cudart.cudaStreamEndCapture(stream))
self.cuda_graph_instance = cuassert(cudart.cudaGraphInstantiate(self.graph, 0))
else:
noerror = self.context.execute_async_v3(stream) # type: ignore
noerror = self.context.execute_async_v3(stream) # type: ignore[union-attr]
if not noerror:
raise ValueError("ERROR: inference failed.")