From 6b63893348fc302fdf6fc7945c6ad240f7a57f9b Mon Sep 17 00:00:00 2001 From: Seunghoon Lee Date: Sun, 9 Jul 2023 15:47:38 +0900 Subject: [PATCH] Fix terminal hang up with TI on DirectML. --- modules/dml/hijack/__init__.py | 1 + modules/dml/hijack/transformers.py | 22 ++++++++++++++++++++++ 2 files changed, 23 insertions(+) create mode 100644 modules/dml/hijack/transformers.py diff --git a/modules/dml/hijack/__init__.py b/modules/dml/hijack/__init__.py index 7e3424775..8a6462e2f 100644 --- a/modules/dml/hijack/__init__.py +++ b/modules/dml/hijack/__init__.py @@ -4,3 +4,4 @@ import modules.dml.hijack.torch import modules.dml.hijack.realesrgan_model import modules.dml.hijack.plms import modules.dml.hijack.diffusers +import modules.dml.hijack.transformers diff --git a/modules/dml/hijack/transformers.py b/modules/dml/hijack/transformers.py new file mode 100644 index 000000000..712d5281e --- /dev/null +++ b/modules/dml/hijack/transformers.py @@ -0,0 +1,22 @@ +import torch +import transformers.models.clip.modeling_clip + +# Copied from transformers.models.bart.modeling_bart._make_causal_mask +def _make_causal_mask( + input_ids_shape: torch.Size, dtype: torch.dtype, device: torch.device, past_key_values_length: int = 0 +): + """ + Make causal mask used for bi-directional self-attention. + """ + bsz, tgt_len = input_ids_shape + min = torch.tensor(torch.finfo(dtype).min, device="cpu") + mask = torch.full((tgt_len, tgt_len), min, device=device) # https://discord.com/channels/1101998836328697867/1127441997184122920 + mask_cond = torch.arange(mask.size(-1), device=device) + mask.masked_fill_(mask_cond < (mask_cond + 1).view(mask.size(-1), 1), 0) + mask = mask.to(dtype) + + if past_key_values_length > 0: + mask = torch.cat([torch.zeros(tgt_len, past_key_values_length, dtype=dtype, device=device), mask], dim=-1) + return mask[None, None, :, :].expand(bsz, 1, tgt_len, tgt_len + past_key_values_length) + +transformers.models.clip.modeling_clip._make_causal_mask = _make_causal_mask