LatentSync 1.5
This commit is contained in:
+142
-153
@@ -1,153 +1,142 @@
|
||||
# Copyright (c) 2024 Bytedance Ltd. and/or its affiliates
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import os
|
||||
import numpy as np
|
||||
from torch.utils.data import Dataset
|
||||
import torch
|
||||
import random
|
||||
from ..utils.util import gather_video_paths_recursively
|
||||
from ..utils.image_processor import ImageProcessor
|
||||
from ..utils.audio import melspectrogram
|
||||
import math
|
||||
|
||||
from decord import AudioReader, VideoReader, cpu
|
||||
|
||||
|
||||
class SyncNetDataset(Dataset):
|
||||
def __init__(self, data_dir: str, fileslist: str, config):
|
||||
if fileslist != "":
|
||||
with open(fileslist) as file:
|
||||
self.video_paths = [line.rstrip() for line in file]
|
||||
elif data_dir != "":
|
||||
self.video_paths = gather_video_paths_recursively(data_dir)
|
||||
else:
|
||||
raise ValueError("data_dir and fileslist cannot be both empty")
|
||||
|
||||
self.resolution = config.data.resolution
|
||||
self.num_frames = config.data.num_frames
|
||||
|
||||
self.mel_window_length = math.ceil(self.num_frames / 5 * 16)
|
||||
|
||||
self.audio_sample_rate = config.data.audio_sample_rate
|
||||
self.video_fps = config.data.video_fps
|
||||
self.audio_samples_length = int(
|
||||
config.data.audio_sample_rate // config.data.video_fps * config.data.num_frames
|
||||
)
|
||||
self.image_processor = ImageProcessor(resolution=config.data.resolution, mask="half")
|
||||
self.audio_cache_dir = config.data.audio_cache_dir
|
||||
os.makedirs(self.audio_cache_dir, exist_ok=True)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.video_paths)
|
||||
|
||||
def read_audio(self, video_path: str):
|
||||
ar = AudioReader(video_path, ctx=cpu(self.worker_id), sample_rate=self.audio_sample_rate)
|
||||
original_mel = melspectrogram(ar[:].asnumpy().squeeze(0))
|
||||
return torch.from_numpy(original_mel)
|
||||
|
||||
def crop_audio_window(self, original_mel, start_index):
|
||||
start_idx = int(80.0 * (start_index / float(self.video_fps)))
|
||||
end_idx = start_idx + self.mel_window_length
|
||||
return original_mel[:, start_idx:end_idx].unsqueeze(0)
|
||||
|
||||
def get_frames(self, video_reader: VideoReader):
|
||||
total_num_frames = len(video_reader)
|
||||
|
||||
start_idx = random.randint(0, total_num_frames - self.num_frames)
|
||||
frames_index = np.arange(start_idx, start_idx + self.num_frames, dtype=int)
|
||||
|
||||
while True:
|
||||
wrong_start_idx = random.randint(0, total_num_frames - self.num_frames)
|
||||
# wrong_start_idx = random.randint(
|
||||
# max(0, start_idx - 25), min(total_num_frames - self.num_frames, start_idx + 25)
|
||||
# )
|
||||
if wrong_start_idx == start_idx:
|
||||
continue
|
||||
# if wrong_start_idx >= start_idx - self.num_frames and wrong_start_idx <= start_idx + self.num_frames:
|
||||
# continue
|
||||
wrong_frames_index = np.arange(wrong_start_idx, wrong_start_idx + self.num_frames, dtype=int)
|
||||
break
|
||||
|
||||
frames = video_reader.get_batch(frames_index).asnumpy()
|
||||
wrong_frames = video_reader.get_batch(wrong_frames_index).asnumpy()
|
||||
|
||||
return frames, wrong_frames, start_idx
|
||||
|
||||
def worker_init_fn(self, worker_id):
|
||||
# Initialize the face mesh object in each worker process,
|
||||
# because the face mesh object cannot be called in subprocesses
|
||||
self.worker_id = worker_id
|
||||
# setattr(self, f"image_processor_{worker_id}", ImageProcessor(self.resolution, self.mask))
|
||||
|
||||
def __getitem__(self, idx):
|
||||
# image_processor = getattr(self, f"image_processor_{self.worker_id}")
|
||||
while True:
|
||||
try:
|
||||
idx = random.randint(0, len(self) - 1)
|
||||
|
||||
# Get video file path
|
||||
video_path = self.video_paths[idx]
|
||||
|
||||
vr = VideoReader(video_path, ctx=cpu(self.worker_id))
|
||||
|
||||
if len(vr) < 2 * self.num_frames:
|
||||
continue
|
||||
|
||||
frames, wrong_frames, start_idx = self.get_frames(vr)
|
||||
|
||||
mel_cache_path = os.path.join(
|
||||
self.audio_cache_dir, os.path.basename(video_path).replace(".mp4", "_mel.pt")
|
||||
)
|
||||
|
||||
if os.path.isfile(mel_cache_path):
|
||||
try:
|
||||
original_mel = torch.load(mel_cache_path)
|
||||
except Exception as e:
|
||||
print(f"{type(e).__name__} - {e} - {mel_cache_path}")
|
||||
os.remove(mel_cache_path)
|
||||
original_mel = self.read_audio(video_path)
|
||||
torch.save(original_mel, mel_cache_path)
|
||||
else:
|
||||
original_mel = self.read_audio(video_path)
|
||||
torch.save(original_mel, mel_cache_path)
|
||||
|
||||
mel = self.crop_audio_window(original_mel, start_idx)
|
||||
|
||||
if mel.shape[-1] != self.mel_window_length:
|
||||
continue
|
||||
|
||||
if random.choice([True, False]):
|
||||
y = torch.ones(1).float()
|
||||
chosen_frames = frames
|
||||
else:
|
||||
y = torch.zeros(1).float()
|
||||
chosen_frames = wrong_frames
|
||||
|
||||
chosen_frames = self.image_processor.process_images(chosen_frames)
|
||||
# chosen_frames, _, _ = image_processor.prepare_masks_and_masked_images(
|
||||
# chosen_frames, affine_transform=True
|
||||
# )
|
||||
|
||||
vr.seek(0) # avoid memory leak
|
||||
break
|
||||
|
||||
except Exception as e: # Handle the exception of face not detcted
|
||||
print(f"{type(e).__name__} - {e} - {video_path}")
|
||||
if "vr" in locals():
|
||||
vr.seek(0) # avoid memory leak
|
||||
|
||||
sample = dict(frames=chosen_frames, audio_samples=mel, y=y)
|
||||
|
||||
return sample
|
||||
# Copyright (c) 2024 Bytedance Ltd. and/or its affiliates
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import os
|
||||
import numpy as np
|
||||
from torch.utils.data import Dataset
|
||||
import torch
|
||||
import random
|
||||
from ..utils.util import gather_video_paths_recursively
|
||||
from ..utils.image_processor import ImageProcessor
|
||||
from ..utils.audio import melspectrogram
|
||||
import math
|
||||
|
||||
from decord import AudioReader, VideoReader, cpu
|
||||
|
||||
|
||||
class SyncNetDataset(Dataset):
|
||||
def __init__(self, data_dir: str, fileslist: str, config):
|
||||
if fileslist != "":
|
||||
with open(fileslist) as file:
|
||||
self.video_paths = [line.rstrip() for line in file]
|
||||
elif data_dir != "":
|
||||
self.video_paths = gather_video_paths_recursively(data_dir)
|
||||
else:
|
||||
raise ValueError("data_dir and fileslist cannot be both empty")
|
||||
|
||||
self.resolution = config.data.resolution
|
||||
self.num_frames = config.data.num_frames
|
||||
|
||||
self.mel_window_length = math.ceil(self.num_frames / 5 * 16)
|
||||
|
||||
self.audio_sample_rate = config.data.audio_sample_rate
|
||||
self.video_fps = config.data.video_fps
|
||||
self.image_processor = ImageProcessor(resolution=config.data.resolution, mask="half")
|
||||
self.audio_mel_cache_dir = config.data.audio_mel_cache_dir
|
||||
os.makedirs(self.audio_mel_cache_dir, exist_ok=True)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.video_paths)
|
||||
|
||||
def read_audio(self, video_path: str):
|
||||
ar = AudioReader(video_path, ctx=cpu(self.worker_id), sample_rate=self.audio_sample_rate)
|
||||
original_mel = melspectrogram(ar[:].asnumpy().squeeze(0))
|
||||
return torch.from_numpy(original_mel)
|
||||
|
||||
def crop_audio_window(self, original_mel, start_index):
|
||||
start_idx = int(80.0 * (start_index / float(self.video_fps)))
|
||||
end_idx = start_idx + self.mel_window_length
|
||||
return original_mel[:, start_idx:end_idx].unsqueeze(0)
|
||||
|
||||
def get_frames(self, video_reader: VideoReader):
|
||||
total_num_frames = len(video_reader)
|
||||
|
||||
start_idx = random.randint(0, total_num_frames - self.num_frames)
|
||||
frames_index = np.arange(start_idx, start_idx + self.num_frames, dtype=int)
|
||||
|
||||
while True:
|
||||
wrong_start_idx = random.randint(0, total_num_frames - self.num_frames)
|
||||
if wrong_start_idx == start_idx:
|
||||
continue
|
||||
wrong_frames_index = np.arange(wrong_start_idx, wrong_start_idx + self.num_frames, dtype=int)
|
||||
break
|
||||
|
||||
frames = video_reader.get_batch(frames_index).asnumpy()
|
||||
wrong_frames = video_reader.get_batch(wrong_frames_index).asnumpy()
|
||||
|
||||
return frames, wrong_frames, start_idx
|
||||
|
||||
def worker_init_fn(self, worker_id):
|
||||
# Initialize the face mesh object in each worker process,
|
||||
# because the face mesh object cannot be called in subprocesses
|
||||
self.worker_id = worker_id
|
||||
# setattr(self, f"image_processor_{worker_id}", ImageProcessor(self.resolution, self.mask))
|
||||
|
||||
def __getitem__(self, idx):
|
||||
# image_processor = getattr(self, f"image_processor_{self.worker_id}")
|
||||
while True:
|
||||
try:
|
||||
idx = random.randint(0, len(self) - 1)
|
||||
|
||||
# Get video file path
|
||||
video_path = self.video_paths[idx]
|
||||
|
||||
vr = VideoReader(video_path, ctx=cpu(self.worker_id))
|
||||
|
||||
if len(vr) < 2 * self.num_frames:
|
||||
continue
|
||||
|
||||
frames, wrong_frames, start_idx = self.get_frames(vr)
|
||||
|
||||
mel_cache_path = os.path.join(
|
||||
self.audio_mel_cache_dir, os.path.basename(video_path).replace(".mp4", "_mel.pt")
|
||||
)
|
||||
|
||||
if os.path.isfile(mel_cache_path):
|
||||
try:
|
||||
original_mel = torch.load(mel_cache_path, weights_only=True)
|
||||
except Exception as e:
|
||||
print(f"{type(e).__name__} - {e} - {mel_cache_path}")
|
||||
os.remove(mel_cache_path)
|
||||
original_mel = self.read_audio(video_path)
|
||||
torch.save(original_mel, mel_cache_path)
|
||||
else:
|
||||
original_mel = self.read_audio(video_path)
|
||||
torch.save(original_mel, mel_cache_path)
|
||||
|
||||
mel = self.crop_audio_window(original_mel, start_idx)
|
||||
|
||||
if mel.shape[-1] != self.mel_window_length:
|
||||
continue
|
||||
|
||||
if random.choice([True, False]):
|
||||
y = torch.ones(1).float()
|
||||
chosen_frames = frames
|
||||
else:
|
||||
y = torch.zeros(1).float()
|
||||
chosen_frames = wrong_frames
|
||||
|
||||
chosen_frames = self.image_processor.process_images(chosen_frames)
|
||||
|
||||
vr.seek(0) # avoid memory leak
|
||||
break
|
||||
|
||||
except Exception as e: # Handle the exception of face not detcted
|
||||
print(f"{type(e).__name__} - {e} - {video_path}")
|
||||
if "vr" in locals():
|
||||
vr.seek(0) # avoid memory leak
|
||||
|
||||
sample = dict(frames=chosen_frames, audio_samples=mel, y=y)
|
||||
|
||||
return sample
|
||||
|
||||
+158
-186
@@ -1,186 +1,158 @@
|
||||
# Copyright (c) 2024 Bytedance Ltd. and/or its affiliates
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import os
|
||||
import numpy as np
|
||||
from torch.utils.data import Dataset
|
||||
import torch
|
||||
import random
|
||||
import cv2
|
||||
from ..utils.image_processor import ImageProcessor, load_fixed_mask
|
||||
from ..utils.audio import melspectrogram
|
||||
from decord import AudioReader, VideoReader, cpu
|
||||
|
||||
|
||||
class UNetDataset(Dataset):
|
||||
def __init__(self, train_data_dir: str, config):
|
||||
if config.data.train_fileslist != "":
|
||||
with open(config.data.train_fileslist) as file:
|
||||
self.video_paths = [line.rstrip() for line in file]
|
||||
elif train_data_dir != "":
|
||||
self.video_paths = []
|
||||
for file in os.listdir(train_data_dir):
|
||||
if file.endswith(".mp4"):
|
||||
self.video_paths.append(os.path.join(train_data_dir, file))
|
||||
else:
|
||||
raise ValueError("data_dir and fileslist cannot be both empty")
|
||||
|
||||
self.resolution = config.data.resolution
|
||||
self.num_frames = config.data.num_frames
|
||||
|
||||
if self.num_frames == 16:
|
||||
self.mel_window_length = 52
|
||||
elif self.num_frames == 5:
|
||||
self.mel_window_length = 16
|
||||
else:
|
||||
raise NotImplementedError("Only support 16 and 5 frames now")
|
||||
|
||||
self.audio_sample_rate = config.data.audio_sample_rate
|
||||
self.video_fps = config.data.video_fps
|
||||
self.mask = config.data.mask
|
||||
self.mask_image = load_fixed_mask(self.resolution)
|
||||
self.load_audio_data = config.model.add_audio_layer and config.run.use_syncnet
|
||||
self.audio_cache_dir = config.data.audio_cache_dir
|
||||
os.makedirs(self.audio_cache_dir, exist_ok=True)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.video_paths)
|
||||
|
||||
def read_audio(self, video_path: str):
|
||||
ar = AudioReader(video_path, ctx=cpu(self.worker_id), sample_rate=self.audio_sample_rate)
|
||||
original_mel = melspectrogram(ar[:].asnumpy().squeeze(0))
|
||||
return torch.from_numpy(original_mel)
|
||||
|
||||
def crop_audio_window(self, original_mel, start_index):
|
||||
start_idx = int(80.0 * (start_index / float(self.video_fps)))
|
||||
end_idx = start_idx + self.mel_window_length
|
||||
return original_mel[:, start_idx:end_idx].unsqueeze(0)
|
||||
|
||||
def crop_overlap_audio_window(self, original_mel, start_index):
|
||||
half_num_frames = self.num_frames // 2
|
||||
if start_index - half_num_frames < 0:
|
||||
return None
|
||||
mels = []
|
||||
for i in range(start_index, start_index + self.num_frames):
|
||||
mel = self.crop_audio_window(original_mel, i - half_num_frames)
|
||||
if mel.shape[-1] != self.mel_window_length:
|
||||
return None
|
||||
mels.append(mel)
|
||||
mel_overlap = torch.stack(mels)
|
||||
return mel_overlap
|
||||
|
||||
def get_frames(self, video_reader: VideoReader):
|
||||
total_num_frames = len(video_reader)
|
||||
|
||||
start_idx = random.randint(self.num_frames // 2, total_num_frames - self.num_frames - self.num_frames // 2)
|
||||
frames_index = np.arange(start_idx, start_idx + self.num_frames, dtype=int)
|
||||
|
||||
while True:
|
||||
wrong_start_idx = random.randint(0, total_num_frames - self.num_frames)
|
||||
if wrong_start_idx > start_idx - self.num_frames and wrong_start_idx < start_idx + self.num_frames:
|
||||
continue
|
||||
wrong_frames_index = np.arange(wrong_start_idx, wrong_start_idx + self.num_frames, dtype=int)
|
||||
break
|
||||
|
||||
frames = video_reader.get_batch(frames_index).asnumpy()
|
||||
wrong_frames = video_reader.get_batch(wrong_frames_index).asnumpy()
|
||||
|
||||
return frames, wrong_frames, start_idx
|
||||
|
||||
def worker_init_fn(self, worker_id):
|
||||
# Initialize the face mesh object in each worker process,
|
||||
# because the face mesh object cannot be called in subprocesses
|
||||
self.worker_id = worker_id
|
||||
setattr(
|
||||
self,
|
||||
f"image_processor_{worker_id}",
|
||||
ImageProcessor(self.resolution, self.mask, mask_image=self.mask_image),
|
||||
)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
image_processor = getattr(self, f"image_processor_{self.worker_id}")
|
||||
while True:
|
||||
try:
|
||||
idx = random.randint(0, len(self) - 1)
|
||||
|
||||
# Get video file path
|
||||
video_path = self.video_paths[idx]
|
||||
|
||||
vr = VideoReader(video_path, ctx=cpu(self.worker_id))
|
||||
|
||||
if len(vr) < 3 * self.num_frames:
|
||||
continue
|
||||
|
||||
continuous_frames, ref_frames, start_idx = self.get_frames(vr)
|
||||
|
||||
if self.load_audio_data:
|
||||
mel_cache_path = os.path.join(
|
||||
"/mnt/bn/maliva-gen-ai-v2/chunyu.li/audio_cache/mel_new",
|
||||
os.path.basename(video_path).replace(".mp4", "_mel.pt"),
|
||||
)
|
||||
|
||||
if os.path.isfile(mel_cache_path):
|
||||
try:
|
||||
original_mel = torch.load(mel_cache_path)
|
||||
except Exception as e:
|
||||
print(f"{type(e).__name__} - {e} - {mel_cache_path}")
|
||||
os.remove(mel_cache_path)
|
||||
original_mel = self.read_audio(video_path)
|
||||
torch.save(original_mel, mel_cache_path)
|
||||
else:
|
||||
original_mel = self.read_audio(video_path)
|
||||
torch.save(original_mel, mel_cache_path)
|
||||
|
||||
mel = self.crop_audio_window(original_mel, start_idx)
|
||||
|
||||
if mel.shape[-1] != self.mel_window_length:
|
||||
continue
|
||||
|
||||
# mel_overlap = self.crop_overlap_audio_window(original_mel, start_idx)
|
||||
|
||||
# if mel_overlap is None:
|
||||
# continue
|
||||
mel_overlap = []
|
||||
else:
|
||||
mel = []
|
||||
mel_overlap = []
|
||||
|
||||
gt, masked_gt, mask = image_processor.prepare_masks_and_masked_images(
|
||||
continuous_frames, affine_transform=False
|
||||
)
|
||||
|
||||
if self.mask == "fix_mask":
|
||||
ref, _, _ = image_processor.prepare_masks_and_masked_images(ref_frames, affine_transform=False)
|
||||
else:
|
||||
ref = image_processor.process_images(ref_frames)
|
||||
vr.seek(0) # avoid memory leak
|
||||
break
|
||||
|
||||
except Exception as e: # Handle the exception of face not detcted
|
||||
print(f"{type(e).__name__} - {e} - {video_path}")
|
||||
if "vr" in locals():
|
||||
vr.seek(0) # avoid memory leak
|
||||
|
||||
sample = dict(
|
||||
gt=gt,
|
||||
masked_gt=masked_gt,
|
||||
ref=ref,
|
||||
mel_overlap=mel_overlap,
|
||||
mel=mel,
|
||||
mask=mask,
|
||||
video_path=video_path,
|
||||
start_idx=start_idx,
|
||||
)
|
||||
|
||||
return sample
|
||||
# Copyright (c) 2024 Bytedance Ltd. and/or its affiliates
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import os
|
||||
import math
|
||||
import numpy as np
|
||||
from torch.utils.data import Dataset
|
||||
import torch
|
||||
import random
|
||||
import cv2
|
||||
from ..utils.image_processor import ImageProcessor, load_fixed_mask
|
||||
from ..utils.audio import melspectrogram
|
||||
from decord import AudioReader, VideoReader, cpu
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class UNetDataset(Dataset):
|
||||
def __init__(self, train_data_dir: str, config):
|
||||
if config.data.train_fileslist != "":
|
||||
with open(config.data.train_fileslist) as file:
|
||||
self.video_paths = [line.rstrip() for line in file]
|
||||
elif train_data_dir != "":
|
||||
self.video_paths = []
|
||||
for file in os.listdir(train_data_dir):
|
||||
if file.endswith(".mp4"):
|
||||
self.video_paths.append(os.path.join(train_data_dir, file))
|
||||
else:
|
||||
raise ValueError("data_dir and fileslist cannot be both empty")
|
||||
|
||||
self.resolution = config.data.resolution
|
||||
self.num_frames = config.data.num_frames
|
||||
|
||||
self.mel_window_length = math.ceil(self.num_frames / 5 * 16)
|
||||
|
||||
self.audio_sample_rate = config.data.audio_sample_rate
|
||||
self.video_fps = config.data.video_fps
|
||||
self.mask = config.data.mask
|
||||
self.mask_image = load_fixed_mask(self.resolution, config.data.mask_image_path)
|
||||
self.load_audio_data = config.model.add_audio_layer and config.run.use_syncnet
|
||||
self.audio_mel_cache_dir = config.data.audio_mel_cache_dir
|
||||
os.makedirs(self.audio_mel_cache_dir, exist_ok=True)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.video_paths)
|
||||
|
||||
def read_audio(self, video_path: str):
|
||||
ar = AudioReader(video_path, ctx=cpu(self.worker_id), sample_rate=self.audio_sample_rate)
|
||||
original_mel = melspectrogram(ar[:].asnumpy().squeeze(0))
|
||||
return torch.from_numpy(original_mel)
|
||||
|
||||
def crop_audio_window(self, original_mel, start_index):
|
||||
start_idx = int(80.0 * (start_index / float(self.video_fps)))
|
||||
end_idx = start_idx + self.mel_window_length
|
||||
return original_mel[:, start_idx:end_idx].unsqueeze(0)
|
||||
|
||||
def get_frames(self, video_reader: VideoReader):
|
||||
total_num_frames = len(video_reader)
|
||||
|
||||
start_idx = random.randint(0, total_num_frames - self.num_frames)
|
||||
gt_frames_index = np.arange(start_idx, start_idx + self.num_frames, dtype=int)
|
||||
|
||||
while True:
|
||||
ref_start_idx = random.randint(0, total_num_frames - self.num_frames)
|
||||
if ref_start_idx > start_idx - self.num_frames and ref_start_idx < start_idx + self.num_frames:
|
||||
continue
|
||||
ref_frames_index = np.arange(ref_start_idx, ref_start_idx + self.num_frames, dtype=int)
|
||||
break
|
||||
|
||||
gt_frames = video_reader.get_batch(gt_frames_index).asnumpy()
|
||||
ref_frames = video_reader.get_batch(ref_frames_index).asnumpy()
|
||||
|
||||
return gt_frames, ref_frames, start_idx
|
||||
|
||||
def worker_init_fn(self, worker_id):
|
||||
# Initialize the face mesh object in each worker process,
|
||||
# because the face mesh object cannot be called in subprocesses
|
||||
self.worker_id = worker_id
|
||||
setattr(
|
||||
self,
|
||||
f"image_processor_{worker_id}",
|
||||
ImageProcessor(self.resolution, self.mask, mask_image=self.mask_image),
|
||||
)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
image_processor: ImageProcessor = getattr(self, f"image_processor_{self.worker_id}")
|
||||
while True:
|
||||
try:
|
||||
idx = random.randint(0, len(self) - 1)
|
||||
|
||||
# Get video file path
|
||||
video_path = self.video_paths[idx]
|
||||
|
||||
vr = VideoReader(video_path, ctx=cpu(self.worker_id))
|
||||
|
||||
if len(vr) < 3 * self.num_frames:
|
||||
continue
|
||||
|
||||
gt_frames, ref_frames, start_idx = self.get_frames(vr)
|
||||
|
||||
if self.load_audio_data:
|
||||
mel_cache_path = os.path.join(
|
||||
self.audio_mel_cache_dir, os.path.basename(video_path).replace(".mp4", "_mel.pt")
|
||||
)
|
||||
|
||||
if os.path.isfile(mel_cache_path):
|
||||
try:
|
||||
original_mel = torch.load(mel_cache_path, weights_only=True)
|
||||
except Exception as e:
|
||||
print(f"{type(e).__name__} - {e} - {mel_cache_path}")
|
||||
os.remove(mel_cache_path)
|
||||
original_mel = self.read_audio(video_path)
|
||||
torch.save(original_mel, mel_cache_path)
|
||||
else:
|
||||
original_mel = self.read_audio(video_path)
|
||||
torch.save(original_mel, mel_cache_path)
|
||||
|
||||
mel = self.crop_audio_window(original_mel, start_idx)
|
||||
|
||||
if mel.shape[-1] != self.mel_window_length:
|
||||
continue
|
||||
else:
|
||||
mel = []
|
||||
|
||||
gt_pixel_values, masked_pixel_values, masks = image_processor.prepare_masks_and_masked_images(
|
||||
gt_frames, affine_transform=False
|
||||
) # (f, c, h, w)
|
||||
ref_pixel_values = image_processor.process_images(ref_frames)
|
||||
|
||||
vr.seek(0) # avoid memory leak
|
||||
break
|
||||
|
||||
except Exception as e: # Handle the exception of face not detcted
|
||||
print(f"{type(e).__name__} - {e} - {video_path}")
|
||||
if "vr" in locals():
|
||||
vr.seek(0) # avoid memory leak
|
||||
|
||||
sample = dict(
|
||||
gt_pixel_values=gt_pixel_values,
|
||||
masked_pixel_values=masked_pixel_values,
|
||||
ref_pixel_values=ref_pixel_values,
|
||||
mel=mel,
|
||||
masks=masks,
|
||||
video_path=video_path,
|
||||
start_idx=start_idx,
|
||||
)
|
||||
|
||||
return sample
|
||||
|
||||
+280
-488
@@ -1,488 +1,280 @@
|
||||
# Adapted from https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention.py
|
||||
from dataclasses import dataclass
|
||||
# Removed turtle import as forward methods are defined in classes
|
||||
from typing import Optional
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers import ModelMixin
|
||||
from diffusers.utils import BaseOutput
|
||||
from diffusers.utils.import_utils import is_xformers_available
|
||||
from diffusers.models.attention import Attention as CrossAttention, FeedForward, AdaLayerNorm
|
||||
from einops import rearrange, repeat
|
||||
from .utils import zero_module
|
||||
|
||||
|
||||
@dataclass
|
||||
class Transformer3DModelOutput(BaseOutput):
|
||||
sample: torch.FloatTensor
|
||||
|
||||
|
||||
if is_xformers_available():
|
||||
import xformers
|
||||
import xformers.ops
|
||||
else:
|
||||
xformers = None
|
||||
|
||||
|
||||
class Transformer3DModel(ModelMixin, ConfigMixin):
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
num_attention_heads: int = 16,
|
||||
attention_head_dim: int = 88,
|
||||
in_channels: Optional[int] = None,
|
||||
num_layers: int = 1,
|
||||
dropout: float = 0.0,
|
||||
norm_num_groups: int = 32,
|
||||
cross_attention_dim: Optional[int] = None,
|
||||
attention_bias: bool = False,
|
||||
activation_fn: str = "geglu",
|
||||
num_embeds_ada_norm: Optional[int] = None,
|
||||
use_linear_projection: bool = False,
|
||||
only_cross_attention: bool = False,
|
||||
upcast_attention: bool = False,
|
||||
use_motion_module: bool = False,
|
||||
unet_use_cross_frame_attention=None,
|
||||
unet_use_temporal_attention=None,
|
||||
add_audio_layer=False,
|
||||
audio_condition_method="cross_attn",
|
||||
custom_audio_layer: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
self.use_linear_projection = use_linear_projection
|
||||
self.num_attention_heads = num_attention_heads
|
||||
self.attention_head_dim = attention_head_dim
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
|
||||
# Define input layers
|
||||
self.in_channels = in_channels
|
||||
|
||||
self.norm = torch.nn.GroupNorm(num_groups=norm_num_groups, num_channels=in_channels, eps=1e-6, affine=True)
|
||||
if use_linear_projection:
|
||||
self.proj_in = nn.Linear(in_channels, inner_dim)
|
||||
else:
|
||||
self.proj_in = nn.Conv2d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0)
|
||||
|
||||
if not custom_audio_layer:
|
||||
# Define transformers blocks
|
||||
self.transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
BasicTransformerBlock(
|
||||
inner_dim,
|
||||
num_attention_heads,
|
||||
attention_head_dim,
|
||||
dropout=dropout,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
activation_fn=activation_fn,
|
||||
num_embeds_ada_norm=num_embeds_ada_norm,
|
||||
attention_bias=attention_bias,
|
||||
only_cross_attention=only_cross_attention,
|
||||
upcast_attention=upcast_attention,
|
||||
use_motion_module=use_motion_module,
|
||||
unet_use_cross_frame_attention=unet_use_cross_frame_attention,
|
||||
unet_use_temporal_attention=unet_use_temporal_attention,
|
||||
add_audio_layer=add_audio_layer,
|
||||
custom_audio_layer=custom_audio_layer,
|
||||
audio_condition_method=audio_condition_method,
|
||||
)
|
||||
for d in range(num_layers)
|
||||
]
|
||||
)
|
||||
else:
|
||||
self.transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
AudioTransformerBlock(
|
||||
inner_dim,
|
||||
num_attention_heads,
|
||||
attention_head_dim,
|
||||
dropout=dropout,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
activation_fn=activation_fn,
|
||||
num_embeds_ada_norm=num_embeds_ada_norm,
|
||||
attention_bias=attention_bias,
|
||||
only_cross_attention=only_cross_attention,
|
||||
upcast_attention=upcast_attention,
|
||||
use_motion_module=use_motion_module,
|
||||
unet_use_cross_frame_attention=unet_use_cross_frame_attention,
|
||||
unet_use_temporal_attention=unet_use_temporal_attention,
|
||||
add_audio_layer=add_audio_layer,
|
||||
)
|
||||
for d in range(num_layers)
|
||||
]
|
||||
)
|
||||
|
||||
# 4. Define output layers
|
||||
if use_linear_projection:
|
||||
self.proj_out = nn.Linear(in_channels, inner_dim)
|
||||
else:
|
||||
self.proj_out = nn.Conv2d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0)
|
||||
|
||||
if custom_audio_layer:
|
||||
self.proj_out = zero_module(self.proj_out)
|
||||
|
||||
def forward(self, hidden_states, encoder_hidden_states=None, timestep=None, return_dict: bool = True):
|
||||
# Input
|
||||
assert hidden_states.dim() == 5, f"Expected hidden_states to have ndim=5, but got ndim={hidden_states.dim()}."
|
||||
video_length = hidden_states.shape[2]
|
||||
hidden_states = rearrange(hidden_states, "b c f h w -> (b f) c h w")
|
||||
|
||||
# No need to do this for audio input, because different audio samples are independent
|
||||
# encoder_hidden_states = repeat(encoder_hidden_states, 'b n c -> (b f) n c', f=video_length)
|
||||
|
||||
batch, channel, height, weight = hidden_states.shape
|
||||
residual = hidden_states
|
||||
|
||||
hidden_states = self.norm(hidden_states)
|
||||
if not self.use_linear_projection:
|
||||
hidden_states = self.proj_in(hidden_states)
|
||||
inner_dim = hidden_states.shape[1]
|
||||
hidden_states = hidden_states.permute(0, 2, 3, 1).reshape(batch, height * weight, inner_dim)
|
||||
else:
|
||||
inner_dim = hidden_states.shape[1]
|
||||
hidden_states = hidden_states.permute(0, 2, 3, 1).reshape(batch, height * weight, inner_dim)
|
||||
hidden_states = self.proj_in(hidden_states)
|
||||
|
||||
# Blocks
|
||||
for block in self.transformer_blocks:
|
||||
hidden_states = block(
|
||||
hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep,
|
||||
video_length=video_length,
|
||||
)
|
||||
|
||||
# Output
|
||||
if not self.use_linear_projection:
|
||||
hidden_states = hidden_states.reshape(batch, height, weight, inner_dim).permute(0, 3, 1, 2).contiguous()
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
else:
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
hidden_states = hidden_states.reshape(batch, height, weight, inner_dim).permute(0, 3, 1, 2).contiguous()
|
||||
|
||||
output = hidden_states + residual
|
||||
|
||||
output = rearrange(output, "(b f) c h w -> b c f h w", f=video_length)
|
||||
if not return_dict:
|
||||
return (output,)
|
||||
|
||||
return Transformer3DModelOutput(sample=output)
|
||||
|
||||
|
||||
class BasicTransformerBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
dropout=0.0,
|
||||
cross_attention_dim: Optional[int] = None,
|
||||
activation_fn: str = "geglu",
|
||||
num_embeds_ada_norm: Optional[int] = None,
|
||||
attention_bias: bool = False,
|
||||
only_cross_attention: bool = False,
|
||||
upcast_attention: bool = False,
|
||||
use_motion_module: bool = False,
|
||||
unet_use_cross_frame_attention=None,
|
||||
unet_use_temporal_attention=None,
|
||||
add_audio_layer=False,
|
||||
custom_audio_layer=False,
|
||||
audio_condition_method="cross_attn",
|
||||
):
|
||||
super().__init__()
|
||||
self.only_cross_attention = only_cross_attention
|
||||
self.use_ada_layer_norm = num_embeds_ada_norm is not None
|
||||
self.unet_use_cross_frame_attention = unet_use_cross_frame_attention
|
||||
self.unet_use_temporal_attention = unet_use_temporal_attention
|
||||
self.use_motion_module = use_motion_module
|
||||
self.add_audio_layer = add_audio_layer
|
||||
|
||||
# SC-Attn
|
||||
assert unet_use_cross_frame_attention is not None
|
||||
if unet_use_cross_frame_attention:
|
||||
raise NotImplementedError("SparseCausalAttention2D not implemented yet.")
|
||||
else:
|
||||
self.attn1 = CrossAttention(
|
||||
query_dim=dim,
|
||||
heads=num_attention_heads,
|
||||
dim_head=attention_head_dim,
|
||||
dropout=dropout,
|
||||
bias=attention_bias,
|
||||
upcast_attention=upcast_attention,
|
||||
)
|
||||
self.norm1 = AdaLayerNorm(dim, num_embeds_ada_norm) if self.use_ada_layer_norm else nn.LayerNorm(dim)
|
||||
|
||||
# Cross-Attn
|
||||
if add_audio_layer and audio_condition_method == "cross_attn" and not custom_audio_layer:
|
||||
self.audio_cross_attn = AudioCrossAttn(
|
||||
dim=dim,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=attention_head_dim,
|
||||
dropout=dropout,
|
||||
attention_bias=attention_bias,
|
||||
upcast_attention=upcast_attention,
|
||||
num_embeds_ada_norm=num_embeds_ada_norm,
|
||||
use_ada_layer_norm=self.use_ada_layer_norm,
|
||||
zero_proj_out=False,
|
||||
)
|
||||
else:
|
||||
self.audio_cross_attn = None
|
||||
|
||||
# Feed-forward
|
||||
self.ff = FeedForward(dim, dropout=dropout, activation_fn=activation_fn)
|
||||
self.norm3 = nn.LayerNorm(dim)
|
||||
|
||||
# Temp-Attn
|
||||
assert unet_use_temporal_attention is not None
|
||||
if unet_use_temporal_attention:
|
||||
self.attn_temp = CrossAttention(
|
||||
query_dim=dim,
|
||||
heads=num_attention_heads,
|
||||
dim_head=attention_head_dim,
|
||||
dropout=dropout,
|
||||
bias=attention_bias,
|
||||
upcast_attention=upcast_attention,
|
||||
)
|
||||
nn.init.zeros_(self.attn_temp.to_out[0].weight.data)
|
||||
self.norm_temp = AdaLayerNorm(dim, num_embeds_ada_norm) if self.use_ada_layer_norm else nn.LayerNorm(dim)
|
||||
|
||||
def set_use_memory_efficient_attention_xformers(self, use_memory_efficient_attention_xformers: bool):
|
||||
if not is_xformers_available():
|
||||
print("Here is how to install it")
|
||||
raise ModuleNotFoundError(
|
||||
"Refer to https://github.com/facebookresearch/xformers for more information on how to install"
|
||||
" xformers",
|
||||
name="xformers",
|
||||
)
|
||||
elif not torch.cuda.is_available():
|
||||
raise ValueError(
|
||||
"torch.cuda.is_available() should be True but is False. xformers' memory efficient attention is only"
|
||||
" available for GPU "
|
||||
)
|
||||
else:
|
||||
try:
|
||||
# Make sure we can run the memory efficient attention
|
||||
_ = xformers.ops.memory_efficient_attention(
|
||||
torch.randn((1, 2, 40), device="cuda"),
|
||||
torch.randn((1, 2, 40), device="cuda"),
|
||||
torch.randn((1, 2, 40), device="cuda"),
|
||||
)
|
||||
except Exception as e:
|
||||
raise e
|
||||
self.attn1._use_memory_efficient_attention_xformers = use_memory_efficient_attention_xformers
|
||||
if self.audio_cross_attn is not None:
|
||||
self.audio_cross_attn.attn._use_memory_efficient_attention_xformers = (
|
||||
use_memory_efficient_attention_xformers
|
||||
)
|
||||
# self.attn_temp._use_memory_efficient_attention_xformers = use_memory_efficient_attention_xformers
|
||||
|
||||
def forward(
|
||||
self, hidden_states, encoder_hidden_states=None, timestep=None, attention_mask=None, video_length=None
|
||||
):
|
||||
# SparseCausal-Attention
|
||||
norm_hidden_states = (
|
||||
self.norm1(hidden_states, timestep) if self.use_ada_layer_norm else self.norm1(hidden_states)
|
||||
)
|
||||
|
||||
# if self.only_cross_attention:
|
||||
# hidden_states = (
|
||||
# self.attn1(norm_hidden_states, encoder_hidden_states, attention_mask=attention_mask) + hidden_states
|
||||
# )
|
||||
# else:
|
||||
# hidden_states = self.attn1(norm_hidden_states, attention_mask=attention_mask, video_length=video_length) + hidden_states
|
||||
|
||||
# pdb.set_trace()
|
||||
if self.unet_use_cross_frame_attention:
|
||||
hidden_states = (
|
||||
self.attn1(norm_hidden_states, attention_mask=attention_mask, video_length=video_length)
|
||||
+ hidden_states
|
||||
)
|
||||
else:
|
||||
hidden_states = self.attn1(norm_hidden_states, attention_mask=attention_mask) + hidden_states
|
||||
|
||||
if self.audio_cross_attn is not None and encoder_hidden_states is not None:
|
||||
hidden_states = self.audio_cross_attn(
|
||||
hidden_states, encoder_hidden_states=encoder_hidden_states, attention_mask=attention_mask
|
||||
)
|
||||
|
||||
# Feed-forward
|
||||
hidden_states = self.ff(self.norm3(hidden_states)) + hidden_states
|
||||
|
||||
# Temporal-Attention
|
||||
if self.unet_use_temporal_attention:
|
||||
d = hidden_states.shape[1]
|
||||
hidden_states = rearrange(hidden_states, "(b f) d c -> (b d) f c", f=video_length)
|
||||
norm_hidden_states = (
|
||||
self.norm_temp(hidden_states, timestep) if self.use_ada_layer_norm else self.norm_temp(hidden_states)
|
||||
)
|
||||
hidden_states = self.attn_temp(norm_hidden_states) + hidden_states
|
||||
hidden_states = rearrange(hidden_states, "(b d) f c -> (b f) d c", d=d)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class AudioTransformerBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
dropout=0.0,
|
||||
cross_attention_dim: Optional[int] = None,
|
||||
activation_fn: str = "geglu",
|
||||
num_embeds_ada_norm: Optional[int] = None,
|
||||
attention_bias: bool = False,
|
||||
only_cross_attention: bool = False,
|
||||
upcast_attention: bool = False,
|
||||
use_motion_module: bool = False,
|
||||
unet_use_cross_frame_attention=None,
|
||||
unet_use_temporal_attention=None,
|
||||
add_audio_layer=False,
|
||||
):
|
||||
super().__init__()
|
||||
self.only_cross_attention = only_cross_attention
|
||||
self.use_ada_layer_norm = num_embeds_ada_norm is not None
|
||||
self.unet_use_cross_frame_attention = unet_use_cross_frame_attention
|
||||
self.unet_use_temporal_attention = unet_use_temporal_attention
|
||||
self.use_motion_module = use_motion_module
|
||||
self.add_audio_layer = add_audio_layer
|
||||
|
||||
# SC-Attn
|
||||
assert unet_use_cross_frame_attention is not None
|
||||
if unet_use_cross_frame_attention:
|
||||
raise NotImplementedError("SparseCausalAttention2D not implemented yet.")
|
||||
else:
|
||||
self.attn1 = CrossAttention(
|
||||
query_dim=dim,
|
||||
heads=num_attention_heads,
|
||||
dim_head=attention_head_dim,
|
||||
dropout=dropout,
|
||||
bias=attention_bias,
|
||||
upcast_attention=upcast_attention,
|
||||
)
|
||||
self.norm1 = AdaLayerNorm(dim, num_embeds_ada_norm) if self.use_ada_layer_norm else nn.LayerNorm(dim)
|
||||
|
||||
self.audio_cross_attn = AudioCrossAttn(
|
||||
dim=dim,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=attention_head_dim,
|
||||
dropout=dropout,
|
||||
attention_bias=attention_bias,
|
||||
upcast_attention=upcast_attention,
|
||||
num_embeds_ada_norm=num_embeds_ada_norm,
|
||||
use_ada_layer_norm=self.use_ada_layer_norm,
|
||||
zero_proj_out=False,
|
||||
)
|
||||
|
||||
# Feed-forward
|
||||
self.ff = FeedForward(dim, dropout=dropout, activation_fn=activation_fn)
|
||||
self.norm3 = nn.LayerNorm(dim)
|
||||
|
||||
def set_use_memory_efficient_attention_xformers(self, use_memory_efficient_attention_xformers: bool):
|
||||
if not is_xformers_available():
|
||||
print("Here is how to install it")
|
||||
raise ModuleNotFoundError(
|
||||
"Refer to https://github.com/facebookresearch/xformers for more information on how to install"
|
||||
" xformers",
|
||||
name="xformers",
|
||||
)
|
||||
elif not torch.cuda.is_available():
|
||||
raise ValueError(
|
||||
"torch.cuda.is_available() should be True but is False. xformers' memory efficient attention is only"
|
||||
" available for GPU "
|
||||
)
|
||||
else:
|
||||
try:
|
||||
# Make sure we can run the memory efficient attention
|
||||
_ = xformers.ops.memory_efficient_attention(
|
||||
torch.randn((1, 2, 40), device="cuda"),
|
||||
torch.randn((1, 2, 40), device="cuda"),
|
||||
torch.randn((1, 2, 40), device="cuda"),
|
||||
)
|
||||
except Exception as e:
|
||||
raise e
|
||||
self.attn1._use_memory_efficient_attention_xformers = use_memory_efficient_attention_xformers
|
||||
if self.audio_cross_attn is not None:
|
||||
self.audio_cross_attn.attn._use_memory_efficient_attention_xformers = (
|
||||
use_memory_efficient_attention_xformers
|
||||
)
|
||||
# self.attn_temp._use_memory_efficient_attention_xformers = use_memory_efficient_attention_xformers
|
||||
|
||||
def forward(
|
||||
self, hidden_states, encoder_hidden_states=None, timestep=None, attention_mask=None, video_length=None
|
||||
):
|
||||
# SparseCausal-Attention
|
||||
norm_hidden_states = (
|
||||
self.norm1(hidden_states, timestep) if self.use_ada_layer_norm else self.norm1(hidden_states)
|
||||
)
|
||||
|
||||
# pdb.set_trace()
|
||||
if self.unet_use_cross_frame_attention:
|
||||
hidden_states = (
|
||||
self.attn1(norm_hidden_states, attention_mask=attention_mask, video_length=video_length)
|
||||
+ hidden_states
|
||||
)
|
||||
else:
|
||||
hidden_states = self.attn1(norm_hidden_states, attention_mask=attention_mask) + hidden_states
|
||||
|
||||
if self.audio_cross_attn is not None and encoder_hidden_states is not None:
|
||||
hidden_states = self.audio_cross_attn(
|
||||
hidden_states, encoder_hidden_states=encoder_hidden_states, attention_mask=attention_mask
|
||||
)
|
||||
|
||||
# Feed-forward
|
||||
hidden_states = self.ff(self.norm3(hidden_states)) + hidden_states
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class AudioCrossAttn(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
cross_attention_dim,
|
||||
num_attention_heads,
|
||||
attention_head_dim,
|
||||
dropout,
|
||||
attention_bias,
|
||||
upcast_attention,
|
||||
num_embeds_ada_norm,
|
||||
use_ada_layer_norm,
|
||||
zero_proj_out=False,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.norm = AdaLayerNorm(dim, num_embeds_ada_norm) if use_ada_layer_norm else nn.LayerNorm(dim)
|
||||
self.attn = CrossAttention(
|
||||
query_dim=dim,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
heads=num_attention_heads,
|
||||
dim_head=attention_head_dim,
|
||||
dropout=dropout,
|
||||
bias=attention_bias,
|
||||
upcast_attention=upcast_attention,
|
||||
)
|
||||
|
||||
if zero_proj_out:
|
||||
self.proj_out = zero_module(nn.Linear(dim, dim))
|
||||
|
||||
self.zero_proj_out = zero_proj_out
|
||||
self.use_ada_layer_norm = use_ada_layer_norm
|
||||
|
||||
def forward(self, hidden_states, encoder_hidden_states=None, timestep=None, attention_mask=None):
|
||||
previous_hidden_states = hidden_states
|
||||
hidden_states = self.norm(hidden_states, timestep) if self.use_ada_layer_norm else self.norm(hidden_states)
|
||||
|
||||
if encoder_hidden_states.dim() == 4:
|
||||
encoder_hidden_states = rearrange(encoder_hidden_states, "b f n d -> (b f) n d")
|
||||
|
||||
hidden_states = self.attn(
|
||||
hidden_states, encoder_hidden_states=encoder_hidden_states, attention_mask=attention_mask
|
||||
)
|
||||
|
||||
if self.zero_proj_out:
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
return hidden_states + previous_hidden_states
|
||||
# Adapted from https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention.py
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.models import ModelMixin
|
||||
from diffusers.utils import BaseOutput
|
||||
from diffusers.models.attention import FeedForward, AdaLayerNorm
|
||||
|
||||
from einops import rearrange, repeat
|
||||
|
||||
|
||||
@dataclass
|
||||
class Transformer3DModelOutput(BaseOutput):
|
||||
sample: torch.FloatTensor
|
||||
|
||||
|
||||
class Transformer3DModel(ModelMixin, ConfigMixin):
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
num_attention_heads: int = 16,
|
||||
attention_head_dim: int = 88,
|
||||
in_channels: Optional[int] = None,
|
||||
num_layers: int = 1,
|
||||
dropout: float = 0.0,
|
||||
norm_num_groups: int = 32,
|
||||
cross_attention_dim: Optional[int] = None,
|
||||
attention_bias: bool = False,
|
||||
activation_fn: str = "geglu",
|
||||
num_embeds_ada_norm: Optional[int] = None,
|
||||
use_linear_projection: bool = False,
|
||||
only_cross_attention: bool = False,
|
||||
upcast_attention: bool = False,
|
||||
add_audio_layer=False,
|
||||
):
|
||||
super().__init__()
|
||||
self.use_linear_projection = use_linear_projection
|
||||
self.num_attention_heads = num_attention_heads
|
||||
self.attention_head_dim = attention_head_dim
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
|
||||
# Define input layers
|
||||
self.in_channels = in_channels
|
||||
|
||||
self.norm = torch.nn.GroupNorm(num_groups=norm_num_groups, num_channels=in_channels, eps=1e-6, affine=True)
|
||||
if use_linear_projection:
|
||||
self.proj_in = nn.Linear(in_channels, inner_dim)
|
||||
else:
|
||||
self.proj_in = nn.Conv2d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0)
|
||||
|
||||
# Define transformers blocks
|
||||
self.transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
BasicTransformerBlock(
|
||||
inner_dim,
|
||||
num_attention_heads,
|
||||
attention_head_dim,
|
||||
dropout=dropout,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
activation_fn=activation_fn,
|
||||
num_embeds_ada_norm=num_embeds_ada_norm,
|
||||
attention_bias=attention_bias,
|
||||
upcast_attention=upcast_attention,
|
||||
add_audio_layer=add_audio_layer,
|
||||
)
|
||||
for d in range(num_layers)
|
||||
]
|
||||
)
|
||||
|
||||
# Define output layers
|
||||
if use_linear_projection:
|
||||
self.proj_out = nn.Linear(in_channels, inner_dim)
|
||||
else:
|
||||
self.proj_out = nn.Conv2d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0)
|
||||
|
||||
def forward(self, hidden_states, encoder_hidden_states=None, timestep=None, return_dict: bool = True):
|
||||
# Input
|
||||
assert hidden_states.dim() == 5, f"Expected hidden_states to have ndim=5, but got ndim={hidden_states.dim()}."
|
||||
video_length = hidden_states.shape[2]
|
||||
hidden_states = rearrange(hidden_states, "b c f h w -> (b f) c h w")
|
||||
|
||||
batch, channel, height, weight = hidden_states.shape
|
||||
residual = hidden_states
|
||||
|
||||
hidden_states = self.norm(hidden_states)
|
||||
if not self.use_linear_projection:
|
||||
hidden_states = self.proj_in(hidden_states)
|
||||
inner_dim = hidden_states.shape[1]
|
||||
hidden_states = hidden_states.permute(0, 2, 3, 1).reshape(batch, height * weight, inner_dim)
|
||||
else:
|
||||
inner_dim = hidden_states.shape[1]
|
||||
hidden_states = hidden_states.permute(0, 2, 3, 1).reshape(batch, height * weight, inner_dim)
|
||||
hidden_states = self.proj_in(hidden_states)
|
||||
|
||||
# Blocks
|
||||
for block in self.transformer_blocks:
|
||||
hidden_states = block(
|
||||
hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep,
|
||||
video_length=video_length,
|
||||
)
|
||||
|
||||
# Output
|
||||
if not self.use_linear_projection:
|
||||
hidden_states = hidden_states.reshape(batch, height, weight, inner_dim).permute(0, 3, 1, 2).contiguous()
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
else:
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
hidden_states = hidden_states.reshape(batch, height, weight, inner_dim).permute(0, 3, 1, 2).contiguous()
|
||||
|
||||
output = hidden_states + residual
|
||||
|
||||
output = rearrange(output, "(b f) c h w -> b c f h w", f=video_length)
|
||||
if not return_dict:
|
||||
return (output,)
|
||||
|
||||
return Transformer3DModelOutput(sample=output)
|
||||
|
||||
|
||||
class BasicTransformerBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
dropout=0.0,
|
||||
cross_attention_dim: Optional[int] = None,
|
||||
activation_fn: str = "geglu",
|
||||
num_embeds_ada_norm: Optional[int] = None,
|
||||
attention_bias: bool = False,
|
||||
upcast_attention: bool = False,
|
||||
add_audio_layer=False,
|
||||
):
|
||||
super().__init__()
|
||||
self.use_ada_layer_norm = num_embeds_ada_norm is not None
|
||||
self.add_audio_layer = add_audio_layer
|
||||
|
||||
self.norm1 = AdaLayerNorm(dim, num_embeds_ada_norm) if self.use_ada_layer_norm else nn.LayerNorm(dim)
|
||||
self.attn1 = Attention(
|
||||
query_dim=dim,
|
||||
heads=num_attention_heads,
|
||||
dim_head=attention_head_dim,
|
||||
dropout=dropout,
|
||||
bias=attention_bias,
|
||||
upcast_attention=upcast_attention,
|
||||
)
|
||||
|
||||
# Cross-attn
|
||||
if add_audio_layer:
|
||||
self.norm2 = AdaLayerNorm(dim, num_embeds_ada_norm) if self.use_ada_layer_norm else nn.LayerNorm(dim)
|
||||
self.attn2 = Attention(
|
||||
query_dim=dim,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
heads=num_attention_heads,
|
||||
dim_head=attention_head_dim,
|
||||
dropout=dropout,
|
||||
bias=attention_bias,
|
||||
upcast_attention=upcast_attention,
|
||||
)
|
||||
else:
|
||||
self.attn2 = None
|
||||
|
||||
# Feed-forward
|
||||
self.ff = FeedForward(dim, dropout=dropout, activation_fn=activation_fn)
|
||||
self.norm3 = nn.LayerNorm(dim)
|
||||
|
||||
def forward(
|
||||
self, hidden_states, encoder_hidden_states=None, timestep=None, attention_mask=None, video_length=None
|
||||
):
|
||||
norm_hidden_states = (
|
||||
self.norm1(hidden_states, timestep) if self.use_ada_layer_norm else self.norm1(hidden_states)
|
||||
)
|
||||
|
||||
hidden_states = self.attn1(norm_hidden_states, attention_mask=attention_mask) + hidden_states
|
||||
|
||||
if self.attn2 is not None and encoder_hidden_states is not None:
|
||||
if encoder_hidden_states.dim() == 4:
|
||||
encoder_hidden_states = rearrange(encoder_hidden_states, "b f s d -> (b f) s d")
|
||||
norm_hidden_states = (
|
||||
self.norm2(hidden_states, timestep) if self.use_ada_layer_norm else self.norm2(hidden_states)
|
||||
)
|
||||
hidden_states = (
|
||||
self.attn2(
|
||||
norm_hidden_states, encoder_hidden_states=encoder_hidden_states, attention_mask=attention_mask
|
||||
)
|
||||
+ hidden_states
|
||||
)
|
||||
|
||||
# Feed-forward
|
||||
hidden_states = self.ff(self.norm3(hidden_states)) + hidden_states
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class Attention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
query_dim: int,
|
||||
cross_attention_dim: Optional[int] = None,
|
||||
heads: int = 8,
|
||||
dim_head: int = 64,
|
||||
dropout: float = 0.0,
|
||||
bias=False,
|
||||
upcast_attention: bool = False,
|
||||
upcast_softmax: bool = False,
|
||||
norm_num_groups: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
inner_dim = dim_head * heads
|
||||
cross_attention_dim = cross_attention_dim if cross_attention_dim is not None else query_dim
|
||||
self.upcast_attention = upcast_attention
|
||||
self.upcast_softmax = upcast_softmax
|
||||
|
||||
self.scale = dim_head**-0.5
|
||||
|
||||
self.heads = heads
|
||||
|
||||
if norm_num_groups is not None:
|
||||
self.group_norm = nn.GroupNorm(num_channels=inner_dim, num_groups=norm_num_groups, eps=1e-5, affine=True)
|
||||
else:
|
||||
self.group_norm = None
|
||||
|
||||
self.to_q = nn.Linear(query_dim, inner_dim, bias=bias)
|
||||
self.to_k = nn.Linear(cross_attention_dim, inner_dim, bias=bias)
|
||||
self.to_v = nn.Linear(cross_attention_dim, inner_dim, bias=bias)
|
||||
|
||||
self.to_out = nn.ModuleList([])
|
||||
self.to_out.append(nn.Linear(inner_dim, query_dim))
|
||||
self.to_out.append(nn.Dropout(dropout))
|
||||
|
||||
def split_heads(self, tensor):
|
||||
batch_size, seq_len, dim = tensor.shape
|
||||
tensor = tensor.reshape(batch_size, seq_len, self.heads, dim // self.heads)
|
||||
tensor = tensor.permute(0, 2, 1, 3)
|
||||
return tensor
|
||||
|
||||
def concat_heads(self, tensor):
|
||||
batch_size, heads, seq_len, head_dim = tensor.shape
|
||||
tensor = tensor.permute(0, 2, 1, 3)
|
||||
tensor = tensor.reshape(batch_size, seq_len, heads * head_dim)
|
||||
return tensor
|
||||
|
||||
def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None):
|
||||
if self.group_norm is not None:
|
||||
hidden_states = self.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
query = self.to_q(hidden_states)
|
||||
query = self.split_heads(query)
|
||||
|
||||
encoder_hidden_states = encoder_hidden_states if encoder_hidden_states is not None else hidden_states
|
||||
key = self.to_k(encoder_hidden_states)
|
||||
value = self.to_v(encoder_hidden_states)
|
||||
|
||||
key = self.split_heads(key)
|
||||
value = self.split_heads(value)
|
||||
|
||||
if attention_mask is not None:
|
||||
if attention_mask.shape[-1] != query.shape[1]:
|
||||
target_length = query.shape[1]
|
||||
attention_mask = F.pad(attention_mask, (0, target_length), value=0.0)
|
||||
attention_mask = attention_mask.repeat_interleave(self.heads, dim=0)
|
||||
|
||||
# Use PyTorch native implementation of FlashAttention-2
|
||||
hidden_states = F.scaled_dot_product_attention(query, key, value, attn_mask=attention_mask)
|
||||
|
||||
hidden_states = self.concat_heads(hidden_states)
|
||||
|
||||
# linear proj
|
||||
hidden_states = self.to_out[0](hidden_states)
|
||||
|
||||
# dropout
|
||||
hidden_states = self.to_out[1](hidden_states)
|
||||
return hidden_states
|
||||
|
||||
+313
-332
@@ -1,332 +1,313 @@
|
||||
# Adapted from https://github.com/guoyww/AnimateDiff/blob/main/animatediff/models/motion_module.py
|
||||
|
||||
# Actually we don't use the motion module in the final version of LatentSync
|
||||
# When we started the project, we used the codebase of AnimateDiff and tried motion module
|
||||
# But the results are poor, and we decied to leave the code here for possible future usage
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers import ModelMixin
|
||||
from diffusers.utils import BaseOutput
|
||||
from diffusers.utils.import_utils import is_xformers_available
|
||||
from diffusers.models.attention import Attention as CrossAttention, FeedForward
|
||||
|
||||
from einops import rearrange, repeat
|
||||
import math
|
||||
from .utils import zero_module
|
||||
|
||||
|
||||
@dataclass
|
||||
class TemporalTransformer3DModelOutput(BaseOutput):
|
||||
sample: torch.FloatTensor
|
||||
|
||||
|
||||
if is_xformers_available():
|
||||
import xformers
|
||||
import xformers.ops
|
||||
else:
|
||||
xformers = None
|
||||
|
||||
|
||||
def get_motion_module(in_channels, motion_module_type: str, motion_module_kwargs: dict):
|
||||
if motion_module_type == "Vanilla":
|
||||
return VanillaTemporalModule(
|
||||
in_channels=in_channels,
|
||||
**motion_module_kwargs,
|
||||
)
|
||||
else:
|
||||
raise ValueError
|
||||
|
||||
|
||||
class VanillaTemporalModule(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
num_attention_heads=8,
|
||||
num_transformer_block=2,
|
||||
attention_block_types=("Temporal_Self", "Temporal_Self"),
|
||||
cross_frame_attention_mode=None,
|
||||
temporal_position_encoding=False,
|
||||
temporal_position_encoding_max_len=24,
|
||||
temporal_attention_dim_div=1,
|
||||
zero_initialize=True,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.temporal_transformer = TemporalTransformer3DModel(
|
||||
in_channels=in_channels,
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=in_channels // num_attention_heads // temporal_attention_dim_div,
|
||||
num_layers=num_transformer_block,
|
||||
attention_block_types=attention_block_types,
|
||||
cross_frame_attention_mode=cross_frame_attention_mode,
|
||||
temporal_position_encoding=temporal_position_encoding,
|
||||
temporal_position_encoding_max_len=temporal_position_encoding_max_len,
|
||||
)
|
||||
|
||||
if zero_initialize:
|
||||
self.temporal_transformer.proj_out = zero_module(self.temporal_transformer.proj_out)
|
||||
|
||||
def forward(self, input_tensor, temb, encoder_hidden_states, attention_mask=None, anchor_frame_idx=None):
|
||||
hidden_states = input_tensor
|
||||
hidden_states = self.temporal_transformer(hidden_states, encoder_hidden_states, attention_mask)
|
||||
|
||||
output = hidden_states
|
||||
return output
|
||||
|
||||
|
||||
class TemporalTransformer3DModel(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
num_attention_heads,
|
||||
attention_head_dim,
|
||||
num_layers,
|
||||
attention_block_types=(
|
||||
"Temporal_Self",
|
||||
"Temporal_Self",
|
||||
),
|
||||
dropout=0.0,
|
||||
norm_num_groups=32,
|
||||
cross_attention_dim=768,
|
||||
activation_fn="geglu",
|
||||
attention_bias=False,
|
||||
upcast_attention=False,
|
||||
cross_frame_attention_mode=None,
|
||||
temporal_position_encoding=False,
|
||||
temporal_position_encoding_max_len=24,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
|
||||
self.norm = torch.nn.GroupNorm(num_groups=norm_num_groups, num_channels=in_channels, eps=1e-6, affine=True)
|
||||
self.proj_in = nn.Linear(in_channels, inner_dim)
|
||||
|
||||
self.transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
TemporalTransformerBlock(
|
||||
dim=inner_dim,
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=attention_head_dim,
|
||||
attention_block_types=attention_block_types,
|
||||
dropout=dropout,
|
||||
norm_num_groups=norm_num_groups,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
activation_fn=activation_fn,
|
||||
attention_bias=attention_bias,
|
||||
upcast_attention=upcast_attention,
|
||||
cross_frame_attention_mode=cross_frame_attention_mode,
|
||||
temporal_position_encoding=temporal_position_encoding,
|
||||
temporal_position_encoding_max_len=temporal_position_encoding_max_len,
|
||||
)
|
||||
for d in range(num_layers)
|
||||
]
|
||||
)
|
||||
self.proj_out = nn.Linear(inner_dim, in_channels)
|
||||
|
||||
def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None):
|
||||
assert hidden_states.dim() == 5, f"Expected hidden_states to have ndim=5, but got ndim={hidden_states.dim()}."
|
||||
video_length = hidden_states.shape[2]
|
||||
hidden_states = rearrange(hidden_states, "b c f h w -> (b f) c h w")
|
||||
|
||||
batch, channel, height, weight = hidden_states.shape
|
||||
residual = hidden_states
|
||||
|
||||
hidden_states = self.norm(hidden_states)
|
||||
hidden_states = hidden_states.permute(0, 2, 3, 1).reshape(batch, height * weight, channel)
|
||||
hidden_states = self.proj_in(hidden_states)
|
||||
|
||||
# Transformer Blocks
|
||||
for block in self.transformer_blocks:
|
||||
hidden_states = block(
|
||||
hidden_states, encoder_hidden_states=encoder_hidden_states, video_length=video_length
|
||||
)
|
||||
|
||||
# output
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
hidden_states = hidden_states.reshape(batch, height, weight, channel).permute(0, 3, 1, 2).contiguous()
|
||||
|
||||
output = hidden_states + residual
|
||||
output = rearrange(output, "(b f) c h w -> b c f h w", f=video_length)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
class TemporalTransformerBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
num_attention_heads,
|
||||
attention_head_dim,
|
||||
attention_block_types=(
|
||||
"Temporal_Self",
|
||||
"Temporal_Self",
|
||||
),
|
||||
dropout=0.0,
|
||||
norm_num_groups=32,
|
||||
cross_attention_dim=768,
|
||||
activation_fn="geglu",
|
||||
attention_bias=False,
|
||||
upcast_attention=False,
|
||||
cross_frame_attention_mode=None,
|
||||
temporal_position_encoding=False,
|
||||
temporal_position_encoding_max_len=24,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
attention_blocks = []
|
||||
norms = []
|
||||
|
||||
for block_name in attention_block_types:
|
||||
attention_blocks.append(
|
||||
VersatileAttention(
|
||||
attention_mode=block_name.split("_")[0],
|
||||
cross_attention_dim=cross_attention_dim if block_name.endswith("_Cross") else None,
|
||||
query_dim=dim,
|
||||
heads=num_attention_heads,
|
||||
dim_head=attention_head_dim,
|
||||
dropout=dropout,
|
||||
bias=attention_bias,
|
||||
upcast_attention=upcast_attention,
|
||||
cross_frame_attention_mode=cross_frame_attention_mode,
|
||||
temporal_position_encoding=temporal_position_encoding,
|
||||
temporal_position_encoding_max_len=temporal_position_encoding_max_len,
|
||||
)
|
||||
)
|
||||
norms.append(nn.LayerNorm(dim))
|
||||
|
||||
self.attention_blocks = nn.ModuleList(attention_blocks)
|
||||
self.norms = nn.ModuleList(norms)
|
||||
|
||||
self.ff = FeedForward(dim, dropout=dropout, activation_fn=activation_fn)
|
||||
self.ff_norm = nn.LayerNorm(dim)
|
||||
|
||||
def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None, video_length=None):
|
||||
for attention_block, norm in zip(self.attention_blocks, self.norms):
|
||||
norm_hidden_states = norm(hidden_states)
|
||||
hidden_states = (
|
||||
attention_block(
|
||||
norm_hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states if attention_block.is_cross_attention else None,
|
||||
video_length=video_length,
|
||||
)
|
||||
+ hidden_states
|
||||
)
|
||||
|
||||
hidden_states = self.ff(self.ff_norm(hidden_states)) + hidden_states
|
||||
|
||||
output = hidden_states
|
||||
return output
|
||||
|
||||
|
||||
class PositionalEncoding(nn.Module):
|
||||
def __init__(self, d_model, dropout=0.0, max_len=24):
|
||||
super().__init__()
|
||||
self.dropout = nn.Dropout(p=dropout)
|
||||
position = torch.arange(max_len).unsqueeze(1)
|
||||
div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model))
|
||||
pe = torch.zeros(1, max_len, d_model)
|
||||
pe[0, :, 0::2] = torch.sin(position * div_term)
|
||||
pe[0, :, 1::2] = torch.cos(position * div_term)
|
||||
self.register_buffer("pe", pe)
|
||||
|
||||
def forward(self, x):
|
||||
x = x + self.pe[:, : x.size(1)]
|
||||
return self.dropout(x)
|
||||
|
||||
|
||||
class VersatileAttention(CrossAttention):
|
||||
def __init__(
|
||||
self,
|
||||
attention_mode=None,
|
||||
cross_frame_attention_mode=None,
|
||||
temporal_position_encoding=False,
|
||||
temporal_position_encoding_max_len=24,
|
||||
*args,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(*args, **kwargs)
|
||||
assert attention_mode == "Temporal"
|
||||
|
||||
self.attention_mode = attention_mode
|
||||
self.is_cross_attention = kwargs["cross_attention_dim"] is not None
|
||||
|
||||
self.pos_encoder = (
|
||||
PositionalEncoding(kwargs["query_dim"], dropout=0.0, max_len=temporal_position_encoding_max_len)
|
||||
if (temporal_position_encoding and attention_mode == "Temporal")
|
||||
else None
|
||||
)
|
||||
|
||||
def extra_repr(self):
|
||||
return f"(Module Info) Attention_Mode: {self.attention_mode}, Is_Cross_Attention: {self.is_cross_attention}"
|
||||
|
||||
def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None, video_length=None):
|
||||
batch_size, sequence_length, _ = hidden_states.shape
|
||||
|
||||
if self.attention_mode == "Temporal":
|
||||
d = hidden_states.shape[1]
|
||||
hidden_states = rearrange(hidden_states, "(b f) d c -> (b d) f c", f=video_length)
|
||||
|
||||
if self.pos_encoder is not None:
|
||||
hidden_states = self.pos_encoder(hidden_states)
|
||||
|
||||
encoder_hidden_states = (
|
||||
repeat(encoder_hidden_states, "b n c -> (b d) n c", d=d)
|
||||
if encoder_hidden_states is not None
|
||||
else encoder_hidden_states
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
# encoder_hidden_states = encoder_hidden_states
|
||||
|
||||
if self.group_norm is not None:
|
||||
hidden_states = self.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
query = self.to_q(hidden_states)
|
||||
dim = query.shape[-1]
|
||||
query = self.reshape_heads_to_batch_dim(query)
|
||||
|
||||
if self.added_kv_proj_dim is not None:
|
||||
raise NotImplementedError
|
||||
|
||||
encoder_hidden_states = encoder_hidden_states if encoder_hidden_states is not None else hidden_states
|
||||
key = self.to_k(encoder_hidden_states)
|
||||
value = self.to_v(encoder_hidden_states)
|
||||
|
||||
key = self.reshape_heads_to_batch_dim(key)
|
||||
value = self.reshape_heads_to_batch_dim(value)
|
||||
|
||||
if attention_mask is not None:
|
||||
if attention_mask.shape[-1] != query.shape[1]:
|
||||
target_length = query.shape[1]
|
||||
attention_mask = F.pad(attention_mask, (0, target_length), value=0.0)
|
||||
attention_mask = attention_mask.repeat_interleave(self.heads, dim=0)
|
||||
|
||||
# attention, what we cannot get enough of
|
||||
if self._use_memory_efficient_attention_xformers:
|
||||
hidden_states = self._memory_efficient_attention_xformers(query, key, value, attention_mask)
|
||||
# Some versions of xformers return output in fp32, cast it back to the dtype of the input
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
else:
|
||||
if self._slice_size is None or query.shape[0] // self._slice_size == 1:
|
||||
hidden_states = self._attention(query, key, value, attention_mask)
|
||||
else:
|
||||
hidden_states = self._sliced_attention(query, key, value, sequence_length, dim, attention_mask)
|
||||
|
||||
# linear proj
|
||||
hidden_states = self.to_out[0](hidden_states)
|
||||
|
||||
# dropout
|
||||
hidden_states = self.to_out[1](hidden_states)
|
||||
|
||||
if self.attention_mode == "Temporal":
|
||||
hidden_states = rearrange(hidden_states, "(b d) f c -> (b f) d c", d=d)
|
||||
|
||||
return hidden_states
|
||||
# Adapted from https://github.com/guoyww/AnimateDiff/blob/main/animatediff/models/motion_module.py
|
||||
|
||||
# Actually we don't use the motion module in the final version of LatentSync
|
||||
# When we started the project, we used the codebase of AnimateDiff and tried motion module
|
||||
# But the results are poor, and we decied to leave the code here for possible future usage
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.models import ModelMixin
|
||||
from diffusers.utils import BaseOutput
|
||||
from diffusers.models.attention import FeedForward
|
||||
from .attention import Attention
|
||||
|
||||
from einops import rearrange, repeat
|
||||
import math
|
||||
from .utils import zero_module
|
||||
|
||||
|
||||
@dataclass
|
||||
class TemporalTransformer3DModelOutput(BaseOutput):
|
||||
sample: torch.FloatTensor
|
||||
|
||||
|
||||
def get_motion_module(in_channels, motion_module_type: str, motion_module_kwargs: dict):
|
||||
if motion_module_type == "Vanilla":
|
||||
return VanillaTemporalModule(
|
||||
in_channels=in_channels,
|
||||
**motion_module_kwargs,
|
||||
)
|
||||
else:
|
||||
raise ValueError
|
||||
|
||||
|
||||
class VanillaTemporalModule(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
num_attention_heads=8,
|
||||
num_transformer_block=2,
|
||||
attention_block_types=("Temporal_Self", "Temporal_Self"),
|
||||
cross_frame_attention_mode=None,
|
||||
temporal_position_encoding=False,
|
||||
temporal_position_encoding_max_len=24,
|
||||
temporal_attention_dim_div=1,
|
||||
zero_initialize=True,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.temporal_transformer = TemporalTransformer3DModel(
|
||||
in_channels=in_channels,
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=in_channels // num_attention_heads // temporal_attention_dim_div,
|
||||
num_layers=num_transformer_block,
|
||||
attention_block_types=attention_block_types,
|
||||
cross_frame_attention_mode=cross_frame_attention_mode,
|
||||
temporal_position_encoding=temporal_position_encoding,
|
||||
temporal_position_encoding_max_len=temporal_position_encoding_max_len,
|
||||
)
|
||||
|
||||
if zero_initialize:
|
||||
self.temporal_transformer.proj_out = zero_module(self.temporal_transformer.proj_out)
|
||||
|
||||
def forward(self, input_tensor, temb, encoder_hidden_states, attention_mask=None, anchor_frame_idx=None):
|
||||
hidden_states = input_tensor
|
||||
hidden_states = self.temporal_transformer(hidden_states, encoder_hidden_states, attention_mask)
|
||||
|
||||
output = hidden_states
|
||||
return output
|
||||
|
||||
|
||||
class TemporalTransformer3DModel(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
num_attention_heads,
|
||||
attention_head_dim,
|
||||
num_layers,
|
||||
attention_block_types=(
|
||||
"Temporal_Self",
|
||||
"Temporal_Self",
|
||||
),
|
||||
dropout=0.0,
|
||||
norm_num_groups=32,
|
||||
cross_attention_dim=768,
|
||||
activation_fn="geglu",
|
||||
attention_bias=False,
|
||||
upcast_attention=False,
|
||||
cross_frame_attention_mode=None,
|
||||
temporal_position_encoding=False,
|
||||
temporal_position_encoding_max_len=24,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
|
||||
self.norm = torch.nn.GroupNorm(num_groups=norm_num_groups, num_channels=in_channels, eps=1e-6, affine=True)
|
||||
self.proj_in = nn.Linear(in_channels, inner_dim)
|
||||
|
||||
self.transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
TemporalTransformerBlock(
|
||||
dim=inner_dim,
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=attention_head_dim,
|
||||
attention_block_types=attention_block_types,
|
||||
dropout=dropout,
|
||||
norm_num_groups=norm_num_groups,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
activation_fn=activation_fn,
|
||||
attention_bias=attention_bias,
|
||||
upcast_attention=upcast_attention,
|
||||
cross_frame_attention_mode=cross_frame_attention_mode,
|
||||
temporal_position_encoding=temporal_position_encoding,
|
||||
temporal_position_encoding_max_len=temporal_position_encoding_max_len,
|
||||
)
|
||||
for d in range(num_layers)
|
||||
]
|
||||
)
|
||||
self.proj_out = nn.Linear(inner_dim, in_channels)
|
||||
|
||||
def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None):
|
||||
assert hidden_states.dim() == 5, f"Expected hidden_states to have ndim=5, but got ndim={hidden_states.dim()}."
|
||||
video_length = hidden_states.shape[2]
|
||||
hidden_states = rearrange(hidden_states, "b c f h w -> (b f) c h w")
|
||||
|
||||
batch, channel, height, weight = hidden_states.shape
|
||||
residual = hidden_states
|
||||
|
||||
hidden_states = self.norm(hidden_states)
|
||||
hidden_states = hidden_states.permute(0, 2, 3, 1).reshape(batch, height * weight, channel)
|
||||
hidden_states = self.proj_in(hidden_states)
|
||||
|
||||
# Transformer Blocks
|
||||
for block in self.transformer_blocks:
|
||||
hidden_states = block(
|
||||
hidden_states, encoder_hidden_states=encoder_hidden_states, video_length=video_length
|
||||
)
|
||||
|
||||
# output
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
hidden_states = hidden_states.reshape(batch, height, weight, channel).permute(0, 3, 1, 2).contiguous()
|
||||
|
||||
output = hidden_states + residual
|
||||
output = rearrange(output, "(b f) c h w -> b c f h w", f=video_length)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
class TemporalTransformerBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
num_attention_heads,
|
||||
attention_head_dim,
|
||||
attention_block_types=(
|
||||
"Temporal_Self",
|
||||
"Temporal_Self",
|
||||
),
|
||||
dropout=0.0,
|
||||
norm_num_groups=32,
|
||||
cross_attention_dim=768,
|
||||
activation_fn="geglu",
|
||||
attention_bias=False,
|
||||
upcast_attention=False,
|
||||
cross_frame_attention_mode=None,
|
||||
temporal_position_encoding=False,
|
||||
temporal_position_encoding_max_len=24,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
attention_blocks = []
|
||||
norms = []
|
||||
|
||||
for block_name in attention_block_types:
|
||||
attention_blocks.append(
|
||||
VersatileAttention(
|
||||
attention_mode=block_name.split("_")[0],
|
||||
cross_attention_dim=cross_attention_dim if block_name.endswith("_Cross") else None,
|
||||
query_dim=dim,
|
||||
heads=num_attention_heads,
|
||||
dim_head=attention_head_dim,
|
||||
dropout=dropout,
|
||||
bias=attention_bias,
|
||||
upcast_attention=upcast_attention,
|
||||
cross_frame_attention_mode=cross_frame_attention_mode,
|
||||
temporal_position_encoding=temporal_position_encoding,
|
||||
temporal_position_encoding_max_len=temporal_position_encoding_max_len,
|
||||
)
|
||||
)
|
||||
norms.append(nn.LayerNorm(dim))
|
||||
|
||||
self.attention_blocks = nn.ModuleList(attention_blocks)
|
||||
self.norms = nn.ModuleList(norms)
|
||||
|
||||
self.ff = FeedForward(dim, dropout=dropout, activation_fn=activation_fn)
|
||||
self.ff_norm = nn.LayerNorm(dim)
|
||||
|
||||
def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None, video_length=None):
|
||||
for attention_block, norm in zip(self.attention_blocks, self.norms):
|
||||
norm_hidden_states = norm(hidden_states)
|
||||
hidden_states = (
|
||||
attention_block(
|
||||
norm_hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states if attention_block.is_cross_attention else None,
|
||||
video_length=video_length,
|
||||
)
|
||||
+ hidden_states
|
||||
)
|
||||
|
||||
hidden_states = self.ff(self.ff_norm(hidden_states)) + hidden_states
|
||||
|
||||
output = hidden_states
|
||||
return output
|
||||
|
||||
|
||||
class PositionalEncoding(nn.Module):
|
||||
def __init__(self, d_model, dropout=0.0, max_len=24):
|
||||
super().__init__()
|
||||
self.dropout = nn.Dropout(p=dropout)
|
||||
position = torch.arange(max_len).unsqueeze(1)
|
||||
div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model))
|
||||
pe = torch.zeros(1, max_len, d_model)
|
||||
pe[0, :, 0::2] = torch.sin(position * div_term)
|
||||
pe[0, :, 1::2] = torch.cos(position * div_term)
|
||||
self.register_buffer("pe", pe)
|
||||
|
||||
def forward(self, x):
|
||||
x = x + self.pe[:, : x.size(1)]
|
||||
return self.dropout(x)
|
||||
|
||||
|
||||
class VersatileAttention(Attention):
|
||||
def __init__(
|
||||
self,
|
||||
attention_mode=None,
|
||||
cross_frame_attention_mode=None,
|
||||
temporal_position_encoding=False,
|
||||
temporal_position_encoding_max_len=24,
|
||||
*args,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(*args, **kwargs)
|
||||
assert attention_mode == "Temporal"
|
||||
|
||||
self.attention_mode = attention_mode
|
||||
self.is_cross_attention = kwargs["cross_attention_dim"] is not None
|
||||
|
||||
self.pos_encoder = (
|
||||
PositionalEncoding(kwargs["query_dim"], dropout=0.0, max_len=temporal_position_encoding_max_len)
|
||||
if (temporal_position_encoding and attention_mode == "Temporal")
|
||||
else None
|
||||
)
|
||||
|
||||
def extra_repr(self):
|
||||
return f"(Module Info) Attention_Mode: {self.attention_mode}, Is_Cross_Attention: {self.is_cross_attention}"
|
||||
|
||||
def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None, video_length=None):
|
||||
if self.attention_mode == "Temporal":
|
||||
s = hidden_states.shape[1]
|
||||
hidden_states = rearrange(hidden_states, "(b f) s c -> (b s) f c", f=video_length)
|
||||
|
||||
if self.pos_encoder is not None:
|
||||
hidden_states = self.pos_encoder(hidden_states)
|
||||
|
||||
##### This section will not be executed #####
|
||||
encoder_hidden_states = (
|
||||
repeat(encoder_hidden_states, "b n c -> (b s) n c", s=s)
|
||||
if encoder_hidden_states is not None
|
||||
else encoder_hidden_states
|
||||
)
|
||||
#############################################
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
if self.group_norm is not None:
|
||||
hidden_states = self.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
query = self.to_q(hidden_states)
|
||||
query = self.split_heads(query)
|
||||
|
||||
encoder_hidden_states = encoder_hidden_states if encoder_hidden_states is not None else hidden_states
|
||||
key = self.to_k(encoder_hidden_states)
|
||||
value = self.to_v(encoder_hidden_states)
|
||||
|
||||
key = self.split_heads(key)
|
||||
value = self.split_heads(value)
|
||||
|
||||
if attention_mask is not None:
|
||||
if attention_mask.shape[-1] != query.shape[1]:
|
||||
target_length = query.shape[1]
|
||||
attention_mask = F.pad(attention_mask, (0, target_length), value=0.0)
|
||||
attention_mask = attention_mask.repeat_interleave(self.heads, dim=0)
|
||||
|
||||
# Use PyTorch native implementation of FlashAttention-2
|
||||
hidden_states = F.scaled_dot_product_attention(query, key, value, attn_mask=attention_mask)
|
||||
|
||||
hidden_states = self.concat_heads(hidden_states)
|
||||
|
||||
# linear proj
|
||||
hidden_states = self.to_out[0](hidden_states)
|
||||
|
||||
# dropout
|
||||
hidden_states = self.to_out[1](hidden_states)
|
||||
|
||||
if self.attention_mode == "Temporal":
|
||||
hidden_states = rearrange(hidden_states, "(b s) f c -> (b f) s c", s=s)
|
||||
|
||||
return hidden_states
|
||||
|
||||
+228
-234
@@ -1,234 +1,228 @@
|
||||
# Adapted from https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/resnet.py
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from einops import rearrange
|
||||
|
||||
|
||||
class InflatedConv3d(nn.Conv2d):
|
||||
def forward(self, x):
|
||||
video_length = x.shape[2]
|
||||
|
||||
x = rearrange(x, "b c f h w -> (b f) c h w")
|
||||
x = super().forward(x)
|
||||
x = rearrange(x, "(b f) c h w -> b c f h w", f=video_length)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class InflatedGroupNorm(nn.GroupNorm):
|
||||
def forward(self, x):
|
||||
video_length = x.shape[2]
|
||||
|
||||
x = rearrange(x, "b c f h w -> (b f) c h w")
|
||||
x = super().forward(x)
|
||||
x = rearrange(x, "(b f) c h w -> b c f h w", f=video_length)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class Upsample3D(nn.Module):
|
||||
def __init__(self, channels, use_conv=False, use_conv_transpose=False, out_channels=None, name="conv"):
|
||||
super().__init__()
|
||||
self.channels = channels
|
||||
self.out_channels = out_channels or channels
|
||||
self.use_conv = use_conv
|
||||
self.use_conv_transpose = use_conv_transpose
|
||||
self.name = name
|
||||
|
||||
conv = None
|
||||
if use_conv_transpose:
|
||||
raise NotImplementedError
|
||||
elif use_conv:
|
||||
self.conv = InflatedConv3d(self.channels, self.out_channels, 3, padding=1)
|
||||
|
||||
def forward(self, hidden_states, output_size=None):
|
||||
assert hidden_states.shape[1] == self.channels
|
||||
|
||||
if self.use_conv_transpose:
|
||||
raise NotImplementedError
|
||||
|
||||
# Cast to float32 to as 'upsample_nearest2d_out_frame' op does not support bfloat16
|
||||
dtype = hidden_states.dtype
|
||||
if dtype == torch.bfloat16:
|
||||
hidden_states = hidden_states.to(torch.float32)
|
||||
|
||||
# upsample_nearest_nhwc fails with large batch sizes. see https://github.com/huggingface/diffusers/issues/984
|
||||
if hidden_states.shape[0] >= 64:
|
||||
hidden_states = hidden_states.contiguous()
|
||||
|
||||
# if `output_size` is passed we force the interpolation output
|
||||
# size and do not make use of `scale_factor=2`
|
||||
if output_size is None:
|
||||
hidden_states = F.interpolate(hidden_states, scale_factor=[1.0, 2.0, 2.0], mode="nearest")
|
||||
else:
|
||||
hidden_states = F.interpolate(hidden_states, size=output_size, mode="nearest")
|
||||
|
||||
# If the input is bfloat16, we cast back to bfloat16
|
||||
if dtype == torch.bfloat16:
|
||||
hidden_states = hidden_states.to(dtype)
|
||||
|
||||
# if self.use_conv:
|
||||
# if self.name == "conv":
|
||||
# hidden_states = self.conv(hidden_states)
|
||||
# else:
|
||||
# hidden_states = self.Conv2d_0(hidden_states)
|
||||
hidden_states = self.conv(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class Downsample3D(nn.Module):
|
||||
def __init__(self, channels, use_conv=False, out_channels=None, padding=1, name="conv"):
|
||||
super().__init__()
|
||||
self.channels = channels
|
||||
self.out_channels = out_channels or channels
|
||||
self.use_conv = use_conv
|
||||
self.padding = padding
|
||||
stride = 2
|
||||
self.name = name
|
||||
|
||||
if use_conv:
|
||||
self.conv = InflatedConv3d(self.channels, self.out_channels, 3, stride=stride, padding=padding)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
def forward(self, hidden_states):
|
||||
assert hidden_states.shape[1] == self.channels
|
||||
if self.use_conv and self.padding == 0:
|
||||
raise NotImplementedError
|
||||
|
||||
assert hidden_states.shape[1] == self.channels
|
||||
hidden_states = self.conv(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class ResnetBlock3D(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
in_channels,
|
||||
out_channels=None,
|
||||
conv_shortcut=False,
|
||||
dropout=0.0,
|
||||
temb_channels=512,
|
||||
groups=32,
|
||||
groups_out=None,
|
||||
pre_norm=True,
|
||||
eps=1e-6,
|
||||
non_linearity="swish",
|
||||
time_embedding_norm="default",
|
||||
output_scale_factor=1.0,
|
||||
use_in_shortcut=None,
|
||||
use_inflated_groupnorm=False,
|
||||
):
|
||||
super().__init__()
|
||||
self.pre_norm = pre_norm
|
||||
self.pre_norm = True
|
||||
self.in_channels = in_channels
|
||||
out_channels = in_channels if out_channels is None else out_channels
|
||||
self.out_channels = out_channels
|
||||
self.use_conv_shortcut = conv_shortcut
|
||||
self.time_embedding_norm = time_embedding_norm
|
||||
self.output_scale_factor = output_scale_factor
|
||||
|
||||
if groups_out is None:
|
||||
groups_out = groups
|
||||
|
||||
assert use_inflated_groupnorm != None
|
||||
if use_inflated_groupnorm:
|
||||
self.norm1 = InflatedGroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)
|
||||
else:
|
||||
self.norm1 = torch.nn.GroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)
|
||||
|
||||
self.conv1 = InflatedConv3d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
|
||||
|
||||
if temb_channels is not None:
|
||||
time_emb_proj_out_channels = out_channels
|
||||
# if self.time_embedding_norm == "default":
|
||||
# time_emb_proj_out_channels = out_channels
|
||||
# elif self.time_embedding_norm == "scale_shift":
|
||||
# time_emb_proj_out_channels = out_channels * 2
|
||||
# else:
|
||||
# raise ValueError(f"unknown time_embedding_norm : {self.time_embedding_norm} ")
|
||||
|
||||
self.time_emb_proj = torch.nn.Linear(temb_channels, time_emb_proj_out_channels)
|
||||
else:
|
||||
self.time_emb_proj = None
|
||||
|
||||
if self.time_embedding_norm == "scale_shift":
|
||||
self.double_len_linear = torch.nn.Linear(time_emb_proj_out_channels, 2 * time_emb_proj_out_channels)
|
||||
else:
|
||||
self.double_len_linear = None
|
||||
|
||||
if use_inflated_groupnorm:
|
||||
self.norm2 = InflatedGroupNorm(num_groups=groups_out, num_channels=out_channels, eps=eps, affine=True)
|
||||
else:
|
||||
self.norm2 = torch.nn.GroupNorm(num_groups=groups_out, num_channels=out_channels, eps=eps, affine=True)
|
||||
|
||||
self.dropout = torch.nn.Dropout(dropout)
|
||||
self.conv2 = InflatedConv3d(out_channels, out_channels, kernel_size=3, stride=1, padding=1)
|
||||
|
||||
if non_linearity == "swish":
|
||||
self.nonlinearity = lambda x: F.silu(x)
|
||||
elif non_linearity == "mish":
|
||||
self.nonlinearity = Mish()
|
||||
elif non_linearity == "silu":
|
||||
self.nonlinearity = nn.SiLU()
|
||||
|
||||
self.use_in_shortcut = self.in_channels != self.out_channels if use_in_shortcut is None else use_in_shortcut
|
||||
|
||||
self.conv_shortcut = None
|
||||
if self.use_in_shortcut:
|
||||
self.conv_shortcut = InflatedConv3d(in_channels, out_channels, kernel_size=1, stride=1, padding=0)
|
||||
|
||||
def forward(self, input_tensor, temb):
|
||||
hidden_states = input_tensor
|
||||
|
||||
hidden_states = self.norm1(hidden_states)
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
|
||||
hidden_states = self.conv1(hidden_states)
|
||||
|
||||
if temb is not None:
|
||||
if temb.dim() == 2:
|
||||
# input (1, 1280)
|
||||
temb = self.time_emb_proj(self.nonlinearity(temb))
|
||||
temb = temb[:, :, None, None, None] # unsqueeze
|
||||
else:
|
||||
# input (1, 1280, 16)
|
||||
temb = temb.permute(0, 2, 1)
|
||||
temb = self.time_emb_proj(self.nonlinearity(temb))
|
||||
if self.double_len_linear is not None:
|
||||
temb = self.double_len_linear(self.nonlinearity(temb))
|
||||
temb = temb.permute(0, 2, 1)
|
||||
temb = temb[:, :, :, None, None]
|
||||
|
||||
if temb is not None and self.time_embedding_norm == "default":
|
||||
hidden_states = hidden_states + temb
|
||||
|
||||
hidden_states = self.norm2(hidden_states)
|
||||
|
||||
if temb is not None and self.time_embedding_norm == "scale_shift":
|
||||
scale, shift = torch.chunk(temb, 2, dim=1)
|
||||
hidden_states = hidden_states * (1 + scale) + shift
|
||||
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
|
||||
hidden_states = self.dropout(hidden_states)
|
||||
hidden_states = self.conv2(hidden_states)
|
||||
|
||||
if self.conv_shortcut is not None:
|
||||
input_tensor = self.conv_shortcut(input_tensor)
|
||||
|
||||
output_tensor = (input_tensor + hidden_states) / self.output_scale_factor
|
||||
|
||||
return output_tensor
|
||||
|
||||
|
||||
class Mish(torch.nn.Module):
|
||||
def forward(self, hidden_states):
|
||||
return hidden_states * torch.tanh(torch.nn.functional.softplus(hidden_states))
|
||||
# Adapted from https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/resnet.py
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from einops import rearrange
|
||||
|
||||
|
||||
class InflatedConv3d(nn.Conv2d):
|
||||
def forward(self, x):
|
||||
video_length = x.shape[2]
|
||||
|
||||
x = rearrange(x, "b c f h w -> (b f) c h w")
|
||||
x = super().forward(x)
|
||||
x = rearrange(x, "(b f) c h w -> b c f h w", f=video_length)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class InflatedGroupNorm(nn.GroupNorm):
|
||||
def forward(self, x):
|
||||
video_length = x.shape[2]
|
||||
|
||||
x = rearrange(x, "b c f h w -> (b f) c h w")
|
||||
x = super().forward(x)
|
||||
x = rearrange(x, "(b f) c h w -> b c f h w", f=video_length)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class Upsample3D(nn.Module):
|
||||
def __init__(self, channels, use_conv=False, use_conv_transpose=False, out_channels=None, name="conv"):
|
||||
super().__init__()
|
||||
self.channels = channels
|
||||
self.out_channels = out_channels or channels
|
||||
self.use_conv = use_conv
|
||||
self.use_conv_transpose = use_conv_transpose
|
||||
self.name = name
|
||||
|
||||
conv = None
|
||||
if use_conv_transpose:
|
||||
raise NotImplementedError
|
||||
elif use_conv:
|
||||
self.conv = InflatedConv3d(self.channels, self.out_channels, 3, padding=1)
|
||||
|
||||
def forward(self, hidden_states, output_size=None):
|
||||
assert hidden_states.shape[1] == self.channels
|
||||
|
||||
if self.use_conv_transpose:
|
||||
raise NotImplementedError
|
||||
|
||||
# Cast to float32 to as 'upsample_nearest2d_out_frame' op does not support bfloat16
|
||||
dtype = hidden_states.dtype
|
||||
if dtype == torch.bfloat16:
|
||||
hidden_states = hidden_states.to(torch.float32)
|
||||
|
||||
# upsample_nearest_nhwc fails with large batch sizes. see https://github.com/huggingface/diffusers/issues/984
|
||||
if hidden_states.shape[0] >= 64:
|
||||
hidden_states = hidden_states.contiguous()
|
||||
|
||||
# if `output_size` is passed we force the interpolation output
|
||||
# size and do not make use of `scale_factor=2`
|
||||
if output_size is None:
|
||||
hidden_states = F.interpolate(hidden_states, scale_factor=[1.0, 2.0, 2.0], mode="nearest")
|
||||
else:
|
||||
hidden_states = F.interpolate(hidden_states, size=output_size, mode="nearest")
|
||||
|
||||
# If the input is bfloat16, we cast back to bfloat16
|
||||
if dtype == torch.bfloat16:
|
||||
hidden_states = hidden_states.to(dtype)
|
||||
|
||||
hidden_states = self.conv(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class Downsample3D(nn.Module):
|
||||
def __init__(self, channels, use_conv=False, out_channels=None, padding=1, name="conv"):
|
||||
super().__init__()
|
||||
self.channels = channels
|
||||
self.out_channels = out_channels or channels
|
||||
self.use_conv = use_conv
|
||||
self.padding = padding
|
||||
stride = 2
|
||||
self.name = name
|
||||
|
||||
if use_conv:
|
||||
self.conv = InflatedConv3d(self.channels, self.out_channels, 3, stride=stride, padding=padding)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
def forward(self, hidden_states):
|
||||
assert hidden_states.shape[1] == self.channels
|
||||
if self.use_conv and self.padding == 0:
|
||||
raise NotImplementedError
|
||||
|
||||
assert hidden_states.shape[1] == self.channels
|
||||
hidden_states = self.conv(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class ResnetBlock3D(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
in_channels,
|
||||
out_channels=None,
|
||||
conv_shortcut=False,
|
||||
dropout=0.0,
|
||||
temb_channels=512,
|
||||
groups=32,
|
||||
groups_out=None,
|
||||
pre_norm=True,
|
||||
eps=1e-6,
|
||||
non_linearity="swish",
|
||||
time_embedding_norm="default",
|
||||
output_scale_factor=1.0,
|
||||
use_in_shortcut=None,
|
||||
use_inflated_groupnorm=False,
|
||||
):
|
||||
super().__init__()
|
||||
self.pre_norm = pre_norm
|
||||
self.pre_norm = True
|
||||
self.in_channels = in_channels
|
||||
out_channels = in_channels if out_channels is None else out_channels
|
||||
self.out_channels = out_channels
|
||||
self.use_conv_shortcut = conv_shortcut
|
||||
self.time_embedding_norm = time_embedding_norm
|
||||
self.output_scale_factor = output_scale_factor
|
||||
|
||||
if groups_out is None:
|
||||
groups_out = groups
|
||||
|
||||
assert use_inflated_groupnorm != None
|
||||
if use_inflated_groupnorm:
|
||||
self.norm1 = InflatedGroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)
|
||||
else:
|
||||
self.norm1 = torch.nn.GroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)
|
||||
|
||||
self.conv1 = InflatedConv3d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
|
||||
|
||||
if temb_channels is not None:
|
||||
if self.time_embedding_norm == "default":
|
||||
time_emb_proj_out_channels = out_channels
|
||||
elif self.time_embedding_norm == "scale_shift":
|
||||
time_emb_proj_out_channels = out_channels * 2
|
||||
else:
|
||||
raise ValueError(f"unknown time_embedding_norm : {self.time_embedding_norm} ")
|
||||
|
||||
self.time_emb_proj = torch.nn.Linear(temb_channels, time_emb_proj_out_channels)
|
||||
else:
|
||||
self.time_emb_proj = None
|
||||
|
||||
if self.time_embedding_norm == "scale_shift":
|
||||
self.double_len_linear = torch.nn.Linear(time_emb_proj_out_channels, 2 * time_emb_proj_out_channels)
|
||||
else:
|
||||
self.double_len_linear = None
|
||||
|
||||
if use_inflated_groupnorm:
|
||||
self.norm2 = InflatedGroupNorm(num_groups=groups_out, num_channels=out_channels, eps=eps, affine=True)
|
||||
else:
|
||||
self.norm2 = torch.nn.GroupNorm(num_groups=groups_out, num_channels=out_channels, eps=eps, affine=True)
|
||||
|
||||
self.dropout = torch.nn.Dropout(dropout)
|
||||
self.conv2 = InflatedConv3d(out_channels, out_channels, kernel_size=3, stride=1, padding=1)
|
||||
|
||||
if non_linearity == "swish":
|
||||
self.nonlinearity = lambda x: F.silu(x)
|
||||
elif non_linearity == "mish":
|
||||
self.nonlinearity = Mish()
|
||||
elif non_linearity == "silu":
|
||||
self.nonlinearity = nn.SiLU()
|
||||
|
||||
self.use_in_shortcut = self.in_channels != self.out_channels if use_in_shortcut is None else use_in_shortcut
|
||||
|
||||
self.conv_shortcut = None
|
||||
if self.use_in_shortcut:
|
||||
self.conv_shortcut = InflatedConv3d(in_channels, out_channels, kernel_size=1, stride=1, padding=0)
|
||||
|
||||
def forward(self, input_tensor, temb):
|
||||
hidden_states = input_tensor
|
||||
|
||||
hidden_states = self.norm1(hidden_states)
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
|
||||
hidden_states = self.conv1(hidden_states)
|
||||
|
||||
if temb is not None:
|
||||
if temb.dim() == 2:
|
||||
# input (1, 1280)
|
||||
temb = self.time_emb_proj(self.nonlinearity(temb))
|
||||
temb = temb[:, :, None, None, None] # unsqueeze
|
||||
else:
|
||||
# input (1, 1280, 16)
|
||||
temb = temb.permute(0, 2, 1)
|
||||
temb = self.time_emb_proj(self.nonlinearity(temb))
|
||||
if self.double_len_linear is not None:
|
||||
temb = self.double_len_linear(self.nonlinearity(temb))
|
||||
temb = temb.permute(0, 2, 1)
|
||||
temb = temb[:, :, :, None, None]
|
||||
|
||||
if temb is not None and self.time_embedding_norm == "default":
|
||||
hidden_states = hidden_states + temb
|
||||
|
||||
hidden_states = self.norm2(hidden_states)
|
||||
|
||||
if temb is not None and self.time_embedding_norm == "scale_shift":
|
||||
scale, shift = torch.chunk(temb, 2, dim=1)
|
||||
hidden_states = hidden_states * (1 + scale) + shift
|
||||
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
|
||||
hidden_states = self.dropout(hidden_states)
|
||||
hidden_states = self.conv2(hidden_states)
|
||||
|
||||
if self.conv_shortcut is not None:
|
||||
input_tensor = self.conv_shortcut(input_tensor)
|
||||
|
||||
output_tensor = (input_tensor + hidden_states) / self.output_scale_factor
|
||||
|
||||
return output_tensor
|
||||
|
||||
|
||||
class Mish(torch.nn.Module):
|
||||
def forward(self, hidden_states):
|
||||
return hidden_states * torch.tanh(torch.nn.functional.softplus(hidden_states))
|
||||
|
||||
@@ -0,0 +1,234 @@
|
||||
# Copyright (c) 2024 Bytedance Ltd. and/or its affiliates
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
from einops import rearrange
|
||||
from torch.nn import functional as F
|
||||
from .attention import Attention
|
||||
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from diffusers.models.attention import FeedForward
|
||||
from einops import rearrange
|
||||
|
||||
|
||||
class StableSyncNet(nn.Module):
|
||||
def __init__(self, config, gradient_checkpointing=False):
|
||||
super().__init__()
|
||||
self.audio_encoder = DownEncoder2D(
|
||||
in_channels=config["audio_encoder"]["in_channels"],
|
||||
block_out_channels=config["audio_encoder"]["block_out_channels"],
|
||||
downsample_factors=config["audio_encoder"]["downsample_factors"],
|
||||
dropout=config["audio_encoder"]["dropout"],
|
||||
attn_blocks=config["audio_encoder"]["attn_blocks"],
|
||||
gradient_checkpointing=gradient_checkpointing,
|
||||
)
|
||||
|
||||
self.visual_encoder = DownEncoder2D(
|
||||
in_channels=config["visual_encoder"]["in_channels"],
|
||||
block_out_channels=config["visual_encoder"]["block_out_channels"],
|
||||
downsample_factors=config["visual_encoder"]["downsample_factors"],
|
||||
dropout=config["visual_encoder"]["dropout"],
|
||||
attn_blocks=config["visual_encoder"]["attn_blocks"],
|
||||
gradient_checkpointing=gradient_checkpointing,
|
||||
)
|
||||
|
||||
self.eval()
|
||||
|
||||
def forward(self, image_sequences, audio_sequences):
|
||||
vision_embeds = self.visual_encoder(image_sequences) # (b, c, 1, 1)
|
||||
audio_embeds = self.audio_encoder(audio_sequences) # (b, c, 1, 1)
|
||||
|
||||
vision_embeds = vision_embeds.reshape(vision_embeds.shape[0], -1) # (b, c)
|
||||
audio_embeds = audio_embeds.reshape(audio_embeds.shape[0], -1) # (b, c)
|
||||
|
||||
# Make them unit vectors
|
||||
vision_embeds = F.normalize(vision_embeds, p=2, dim=1)
|
||||
audio_embeds = F.normalize(audio_embeds, p=2, dim=1)
|
||||
|
||||
return vision_embeds, audio_embeds
|
||||
|
||||
|
||||
class ResnetBlock2D(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
dropout: float = 0.0,
|
||||
norm_num_groups: int = 32,
|
||||
eps: float = 1e-6,
|
||||
act_fn: str = "silu",
|
||||
downsample_factor=2,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.norm1 = nn.GroupNorm(num_groups=norm_num_groups, num_channels=in_channels, eps=eps, affine=True)
|
||||
self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
|
||||
|
||||
self.norm2 = nn.GroupNorm(num_groups=norm_num_groups, num_channels=out_channels, eps=eps, affine=True)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1)
|
||||
|
||||
if act_fn == "relu":
|
||||
self.act_fn = nn.ReLU()
|
||||
elif act_fn == "silu":
|
||||
self.act_fn = nn.SiLU()
|
||||
|
||||
if in_channels != out_channels:
|
||||
self.conv_shortcut = nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, padding=0)
|
||||
else:
|
||||
self.conv_shortcut = None
|
||||
|
||||
if isinstance(downsample_factor, list):
|
||||
downsample_factor = tuple(downsample_factor)
|
||||
|
||||
if downsample_factor == 1:
|
||||
self.downsample_conv = None
|
||||
else:
|
||||
self.downsample_conv = nn.Conv2d(
|
||||
out_channels, out_channels, kernel_size=3, stride=downsample_factor, padding=0
|
||||
)
|
||||
self.pad = (0, 1, 0, 1)
|
||||
if isinstance(downsample_factor, tuple):
|
||||
if downsample_factor[0] == 1:
|
||||
self.pad = (0, 1, 1, 1) # The padding order is from back to front
|
||||
elif downsample_factor[1] == 1:
|
||||
self.pad = (1, 1, 0, 1)
|
||||
|
||||
def forward(self, input_tensor):
|
||||
hidden_states = input_tensor
|
||||
|
||||
hidden_states = self.norm1(hidden_states)
|
||||
hidden_states = self.act_fn(hidden_states)
|
||||
|
||||
hidden_states = self.conv1(hidden_states)
|
||||
hidden_states = self.norm2(hidden_states)
|
||||
hidden_states = self.act_fn(hidden_states)
|
||||
|
||||
hidden_states = self.dropout(hidden_states)
|
||||
hidden_states = self.conv2(hidden_states)
|
||||
|
||||
if self.conv_shortcut is not None:
|
||||
input_tensor = self.conv_shortcut(input_tensor)
|
||||
|
||||
hidden_states += input_tensor
|
||||
|
||||
if self.downsample_conv is not None:
|
||||
hidden_states = F.pad(hidden_states, self.pad, mode="constant", value=0)
|
||||
hidden_states = self.downsample_conv(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class AttentionBlock2D(nn.Module):
|
||||
def __init__(self, query_dim, norm_num_groups=32, dropout=0.0):
|
||||
super().__init__()
|
||||
self.norm1 = torch.nn.GroupNorm(num_groups=norm_num_groups, num_channels=query_dim, eps=1e-6, affine=True)
|
||||
self.norm2 = nn.LayerNorm(query_dim)
|
||||
self.norm3 = nn.LayerNorm(query_dim)
|
||||
|
||||
self.ff = FeedForward(query_dim, dropout=dropout, activation_fn="geglu")
|
||||
|
||||
self.conv_in = nn.Conv2d(query_dim, query_dim, kernel_size=1, stride=1, padding=0)
|
||||
self.conv_out = nn.Conv2d(query_dim, query_dim, kernel_size=1, stride=1, padding=0)
|
||||
|
||||
self.attn = Attention(query_dim=query_dim, heads=8, dim_head=query_dim // 8, dropout=dropout, bias=True)
|
||||
|
||||
def forward(self, hidden_states):
|
||||
assert hidden_states.dim() == 4, f"Expected hidden_states to have ndim=4, but got ndim={hidden_states.dim()}."
|
||||
|
||||
batch, channel, height, width = hidden_states.shape
|
||||
residual = hidden_states
|
||||
|
||||
hidden_states = self.norm1(hidden_states)
|
||||
hidden_states = self.conv_in(hidden_states)
|
||||
hidden_states = rearrange(hidden_states, "b c h w -> b (h w) c")
|
||||
|
||||
norm_hidden_states = self.norm2(hidden_states)
|
||||
|
||||
hidden_states = self.attn(norm_hidden_states, attention_mask=None) + hidden_states
|
||||
hidden_states = self.ff(self.norm3(hidden_states)) + hidden_states
|
||||
|
||||
hidden_states = rearrange(hidden_states, "b (h w) c -> b c h w", h=height, w=width)
|
||||
hidden_states = self.conv_out(hidden_states)
|
||||
|
||||
hidden_states = hidden_states + residual
|
||||
return hidden_states
|
||||
|
||||
|
||||
class DownEncoder2D(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels=4 * 16,
|
||||
block_out_channels=[64, 128, 256, 256],
|
||||
downsample_factors=[2, 2, 2, 2],
|
||||
layers_per_block=2,
|
||||
norm_num_groups=32,
|
||||
attn_blocks=[1, 1, 1, 1],
|
||||
dropout: float = 0.0,
|
||||
act_fn="silu",
|
||||
gradient_checkpointing=False,
|
||||
):
|
||||
super().__init__()
|
||||
self.layers_per_block = layers_per_block
|
||||
self.gradient_checkpointing = gradient_checkpointing
|
||||
|
||||
# in
|
||||
self.conv_in = nn.Conv2d(in_channels, block_out_channels[0], kernel_size=3, stride=1, padding=1)
|
||||
|
||||
# down
|
||||
self.down_blocks = nn.ModuleList([])
|
||||
|
||||
output_channels = block_out_channels[0]
|
||||
for i, block_out_channel in enumerate(block_out_channels):
|
||||
input_channels = output_channels
|
||||
output_channels = block_out_channel
|
||||
# is_final_block = i == len(block_out_channels) - 1
|
||||
|
||||
down_block = ResnetBlock2D(
|
||||
in_channels=input_channels,
|
||||
out_channels=output_channels,
|
||||
downsample_factor=downsample_factors[i],
|
||||
norm_num_groups=norm_num_groups,
|
||||
dropout=dropout,
|
||||
act_fn=act_fn,
|
||||
)
|
||||
|
||||
self.down_blocks.append(down_block)
|
||||
|
||||
if attn_blocks[i] == 1:
|
||||
attention_block = AttentionBlock2D(query_dim=output_channels, dropout=dropout)
|
||||
self.down_blocks.append(attention_block)
|
||||
|
||||
# out
|
||||
self.norm_out = nn.GroupNorm(num_channels=block_out_channels[-1], num_groups=norm_num_groups, eps=1e-6)
|
||||
self.act_fn_out = nn.ReLU()
|
||||
|
||||
def forward(self, hidden_states):
|
||||
hidden_states = self.conv_in(hidden_states)
|
||||
|
||||
# down
|
||||
for down_block in self.down_blocks:
|
||||
if self.gradient_checkpointing:
|
||||
hidden_states = torch.utils.checkpoint.checkpoint(down_block, hidden_states, use_reentrant=False)
|
||||
else:
|
||||
hidden_states = down_block(hidden_states)
|
||||
|
||||
# post-process
|
||||
hidden_states = self.norm_out(hidden_states)
|
||||
hidden_states = self.act_fn_out(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
+512
-528
File diff suppressed because it is too large
Load Diff
+777
-903
File diff suppressed because it is too large
Load Diff
+19
-19
@@ -1,19 +1,19 @@
|
||||
# Copyright (c) 2024 Bytedance Ltd. and/or its affiliates
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
def zero_module(module):
|
||||
# Zero out the parameters of a module and return it.
|
||||
for p in module.parameters():
|
||||
p.detach().zero_()
|
||||
return module
|
||||
# Copyright (c) 2024 Bytedance Ltd. and/or its affiliates
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
def zero_module(module):
|
||||
# Zero out the parameters of a module and return it.
|
||||
for p in module.parameters():
|
||||
p.detach().zero_()
|
||||
return module
|
||||
|
||||
@@ -0,0 +1,90 @@
|
||||
# Adapted from https://github.com/primepake/wav2lip_288x288/blob/master/models/syncnetv2.py
|
||||
# The code here is for ablation study.
|
||||
|
||||
from torch import nn
|
||||
from torch.nn import functional as F
|
||||
|
||||
|
||||
class Wav2LipSyncNet(nn.Module):
|
||||
def __init__(self, act_fn="leaky"):
|
||||
super().__init__()
|
||||
|
||||
# input image sequences: (15, 128, 256)
|
||||
self.visual_encoder = nn.Sequential(
|
||||
Conv2d(15, 32, kernel_size=(7, 7), stride=1, padding=3, act_fn=act_fn), # (128, 256)
|
||||
Conv2d(32, 64, kernel_size=5, stride=(1, 2), padding=1, act_fn=act_fn), # (126, 127)
|
||||
Conv2d(64, 64, kernel_size=3, stride=1, padding=1, residual=True, act_fn=act_fn),
|
||||
Conv2d(64, 64, kernel_size=3, stride=1, padding=1, residual=True, act_fn=act_fn),
|
||||
Conv2d(64, 128, kernel_size=3, stride=2, padding=1, act_fn=act_fn), # (63, 64)
|
||||
Conv2d(128, 128, kernel_size=3, stride=1, padding=1, residual=True, act_fn=act_fn),
|
||||
Conv2d(128, 128, kernel_size=3, stride=1, padding=1, residual=True, act_fn=act_fn),
|
||||
Conv2d(128, 128, kernel_size=3, stride=1, padding=1, residual=True, act_fn=act_fn),
|
||||
Conv2d(128, 256, kernel_size=3, stride=3, padding=1, act_fn=act_fn), # (21, 22)
|
||||
Conv2d(256, 256, kernel_size=3, stride=1, padding=1, residual=True, act_fn=act_fn),
|
||||
Conv2d(256, 256, kernel_size=3, stride=1, padding=1, residual=True, act_fn=act_fn),
|
||||
Conv2d(256, 512, kernel_size=3, stride=2, padding=1, act_fn=act_fn), # (11, 11)
|
||||
Conv2d(512, 512, kernel_size=3, stride=1, padding=1, residual=True, act_fn=act_fn),
|
||||
Conv2d(512, 512, kernel_size=3, stride=1, padding=1, residual=True, act_fn=act_fn),
|
||||
Conv2d(512, 1024, kernel_size=3, stride=2, padding=1, act_fn=act_fn), # (6, 6)
|
||||
Conv2d(1024, 1024, kernel_size=3, stride=1, padding=1, residual=True, act_fn=act_fn),
|
||||
Conv2d(1024, 1024, kernel_size=3, stride=1, padding=1, residual=True, act_fn=act_fn),
|
||||
Conv2d(1024, 1024, kernel_size=3, stride=2, padding=1, act_fn="relu"), # (3, 3)
|
||||
Conv2d(1024, 1024, kernel_size=3, stride=1, padding=0, act_fn="relu"), # (1, 1)
|
||||
Conv2d(1024, 1024, kernel_size=1, stride=1, padding=0, act_fn="relu"),
|
||||
)
|
||||
|
||||
# input audio sequences: (1, 80, 16)
|
||||
self.audio_encoder = nn.Sequential(
|
||||
Conv2d(1, 32, kernel_size=3, stride=1, padding=1, act_fn=act_fn),
|
||||
Conv2d(32, 32, kernel_size=3, stride=1, padding=1, residual=True, act_fn=act_fn),
|
||||
Conv2d(32, 32, kernel_size=3, stride=1, padding=1, residual=True, act_fn=act_fn),
|
||||
Conv2d(32, 64, kernel_size=3, stride=(3, 1), padding=1, act_fn=act_fn), # (27, 16)
|
||||
Conv2d(64, 64, kernel_size=3, stride=1, padding=1, residual=True, act_fn=act_fn),
|
||||
Conv2d(64, 64, kernel_size=3, stride=1, padding=1, residual=True, act_fn=act_fn),
|
||||
Conv2d(64, 128, kernel_size=3, stride=3, padding=1, act_fn=act_fn), # (9, 6)
|
||||
Conv2d(128, 128, kernel_size=3, stride=1, padding=1, residual=True, act_fn=act_fn),
|
||||
Conv2d(128, 128, kernel_size=3, stride=1, padding=1, residual=True, act_fn=act_fn),
|
||||
Conv2d(128, 256, kernel_size=3, stride=(3, 2), padding=1, act_fn=act_fn), # (3, 3)
|
||||
Conv2d(256, 256, kernel_size=3, stride=1, padding=1, residual=True, act_fn=act_fn),
|
||||
Conv2d(256, 256, kernel_size=3, stride=1, padding=1, residual=True, act_fn=act_fn),
|
||||
Conv2d(256, 512, kernel_size=3, stride=1, padding=1, act_fn=act_fn),
|
||||
Conv2d(512, 512, kernel_size=3, stride=1, padding=1, residual=True, act_fn=act_fn),
|
||||
Conv2d(512, 512, kernel_size=3, stride=1, padding=1, residual=True, act_fn=act_fn),
|
||||
Conv2d(512, 1024, kernel_size=3, stride=1, padding=0, act_fn="relu"), # (1, 1)
|
||||
Conv2d(1024, 1024, kernel_size=1, stride=1, padding=0, act_fn="relu"),
|
||||
)
|
||||
|
||||
def forward(self, image_sequences, audio_sequences):
|
||||
vision_embeds = self.visual_encoder(image_sequences) # (b, c, 1, 1)
|
||||
audio_embeds = self.audio_encoder(audio_sequences) # (b, c, 1, 1)
|
||||
|
||||
vision_embeds = vision_embeds.reshape(vision_embeds.shape[0], -1) # (b, c)
|
||||
audio_embeds = audio_embeds.reshape(audio_embeds.shape[0], -1) # (b, c)
|
||||
|
||||
# Make them unit vectors
|
||||
vision_embeds = F.normalize(vision_embeds, p=2, dim=1)
|
||||
audio_embeds = F.normalize(audio_embeds, p=2, dim=1)
|
||||
|
||||
return vision_embeds, audio_embeds
|
||||
|
||||
|
||||
class Conv2d(nn.Module):
|
||||
def __init__(self, cin, cout, kernel_size, stride, padding, residual=False, act_fn="relu", *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.conv_block = nn.Sequential(nn.Conv2d(cin, cout, kernel_size, stride, padding), nn.BatchNorm2d(cout))
|
||||
if act_fn == "relu":
|
||||
self.act_fn = nn.ReLU()
|
||||
elif act_fn == "tanh":
|
||||
self.act_fn = nn.Tanh()
|
||||
elif act_fn == "silu":
|
||||
self.act_fn = nn.SiLU()
|
||||
elif act_fn == "leaky":
|
||||
self.act_fn = nn.LeakyReLU(0.2, inplace=True)
|
||||
|
||||
self.residual = residual
|
||||
|
||||
def forward(self, x):
|
||||
out = self.conv_block(x)
|
||||
if self.residual:
|
||||
out += x
|
||||
return self.act_fn(out)
|
||||
@@ -1,522 +1,465 @@
|
||||
# Adapted from https://github.com/guoyww/AnimateDiff/blob/main/animatediff/pipelines/pipeline_animation.py
|
||||
|
||||
import inspect
|
||||
import os
|
||||
import shutil
|
||||
from typing import Callable, List, Optional, Union
|
||||
from dataclasses import dataclass
|
||||
import subprocess
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
|
||||
from diffusers.utils import is_accelerate_available
|
||||
from packaging import version
|
||||
from transformers import AutoProcessor, Wav2Vec2Model
|
||||
|
||||
from diffusers.configuration_utils import FrozenDict
|
||||
from diffusers.models import AutoencoderKL
|
||||
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
|
||||
from diffusers.schedulers import (
|
||||
DDIMScheduler,
|
||||
DPMSolverMultistepScheduler,
|
||||
EulerAncestralDiscreteScheduler,
|
||||
EulerDiscreteScheduler,
|
||||
LMSDiscreteScheduler,
|
||||
PNDMScheduler,
|
||||
)
|
||||
from diffusers.utils import deprecate, logging
|
||||
|
||||
from einops import rearrange
|
||||
|
||||
from ..models.unet import UNet3DConditionModel
|
||||
from ..utils.image_processor import ImageProcessor
|
||||
from ..utils.util import read_video, read_audio, write_video
|
||||
import tqdm
|
||||
import soundfile as sf
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
class LipsyncPipeline(DiffusionPipeline):
|
||||
_optional_components = []
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vae: AutoencoderKL,
|
||||
audio_processor: AutoProcessor,
|
||||
audio_encoder: Wav2Vec2Model,
|
||||
unet: UNet3DConditionModel,
|
||||
scheduler: Union[
|
||||
DDIMScheduler,
|
||||
PNDMScheduler,
|
||||
LMSDiscreteScheduler,
|
||||
EulerDiscreteScheduler,
|
||||
EulerAncestralDiscreteScheduler,
|
||||
DPMSolverMultistepScheduler,
|
||||
],
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
if hasattr(scheduler.config, "steps_offset") and scheduler.config.steps_offset != 1:
|
||||
deprecation_message = (
|
||||
f"The configuration file of this scheduler: {scheduler} is outdated. `steps_offset`"
|
||||
f" should be set to 1 instead of {scheduler.config.steps_offset}. Please make sure "
|
||||
"to update the config accordingly as leaving `steps_offset` might led to incorrect results"
|
||||
" in future versions. If you have downloaded this checkpoint from the Hugging Face Hub,"
|
||||
" it would be very nice if you could open a Pull request for the `scheduler/scheduler_config.json`"
|
||||
" file"
|
||||
)
|
||||
deprecate("steps_offset!=1", "1.0.0", deprecation_message, standard_warn=False)
|
||||
new_config = dict(scheduler.config)
|
||||
new_config["steps_offset"] = 1
|
||||
scheduler._internal_dict = FrozenDict(new_config)
|
||||
|
||||
if hasattr(scheduler.config, "clip_sample") and scheduler.config.clip_sample is True:
|
||||
deprecation_message = (
|
||||
f"The configuration file of this scheduler: {scheduler} has not set the configuration `clip_sample`."
|
||||
" `clip_sample` should be set to False in the configuration file. Please make sure to update the"
|
||||
" config accordingly as not setting `clip_sample` in the config might lead to incorrect results in"
|
||||
" future versions. If you have downloaded this checkpoint from the Hugging Face Hub, it would be very"
|
||||
" nice if you could open a Pull request for the `scheduler/scheduler_config.json` file"
|
||||
)
|
||||
deprecate("clip_sample not set", "1.0.0", deprecation_message, standard_warn=False)
|
||||
new_config = dict(scheduler.config)
|
||||
new_config["clip_sample"] = False
|
||||
scheduler._internal_dict = FrozenDict(new_config)
|
||||
|
||||
is_unet_version_less_0_9_0 = hasattr(unet.config, "_diffusers_version") and version.parse(
|
||||
version.parse(unet.config._diffusers_version).base_version
|
||||
) < version.parse("0.9.0.dev0")
|
||||
is_unet_sample_size_less_64 = hasattr(unet.config, "sample_size") and unet.config.sample_size < 64
|
||||
if is_unet_version_less_0_9_0 and is_unet_sample_size_less_64:
|
||||
deprecation_message = (
|
||||
"The configuration file of the unet has set the default `sample_size` to smaller than"
|
||||
" 64 which seems highly unlikely. If your checkpoint is a fine-tuned version of any of the"
|
||||
" following: \n- CompVis/stable-diffusion-v1-4 \n- CompVis/stable-diffusion-v1-3 \n-"
|
||||
" CompVis/stable-diffusion-v1-2 \n- CompVis/stable-diffusion-v1-1 \n- runwayml/stable-diffusion-v1-5"
|
||||
" \n- runwayml/stable-diffusion-inpainting \n you should change 'sample_size' to 64 in the"
|
||||
" configuration file. Please make sure to update the config accordingly as leaving `sample_size=32`"
|
||||
" in the config might lead to incorrect results in future versions. If you have downloaded this"
|
||||
" checkpoint from the Hugging Face Hub, it would be very nice if you could open a Pull request for"
|
||||
" the `unet/config.json` file"
|
||||
)
|
||||
deprecate("sample_size<64", "1.0.0", deprecation_message, standard_warn=False)
|
||||
new_config = dict(unet.config)
|
||||
new_config["sample_size"] = 64
|
||||
unet._internal_dict = FrozenDict(new_config)
|
||||
|
||||
self.register_modules(
|
||||
vae=vae,
|
||||
audio_processor=audio_processor,
|
||||
audio_encoder=audio_encoder,
|
||||
unet=unet,
|
||||
scheduler=scheduler,
|
||||
)
|
||||
|
||||
self.latent_space = vae is not None
|
||||
if self.latent_space:
|
||||
self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
|
||||
|
||||
self.set_progress_bar_config(desc="Steps")
|
||||
|
||||
def enable_vae_slicing(self):
|
||||
self.vae.enable_slicing()
|
||||
|
||||
def disable_vae_slicing(self):
|
||||
self.vae.disable_slicing()
|
||||
|
||||
def enable_sequential_cpu_offload(self, gpu_id=0):
|
||||
if is_accelerate_available():
|
||||
from accelerate import cpu_offload
|
||||
else:
|
||||
raise ImportError("Please install accelerate via `pip install accelerate`")
|
||||
|
||||
device = torch.device(f"cuda:{gpu_id}")
|
||||
|
||||
for cpu_offloaded_model in [self.unet, self.text_encoder, self.vae]:
|
||||
if cpu_offloaded_model is not None:
|
||||
cpu_offload(cpu_offloaded_model, device)
|
||||
|
||||
@property
|
||||
def _execution_device(self):
|
||||
if self.device != torch.device("meta") or not hasattr(self.unet, "_hf_hook"):
|
||||
return self.device
|
||||
for module in self.unet.modules():
|
||||
if (
|
||||
hasattr(module, "_hf_hook")
|
||||
and hasattr(module._hf_hook, "execution_device")
|
||||
and module._hf_hook.execution_device is not None
|
||||
):
|
||||
return torch.device(module._hf_hook.execution_device)
|
||||
return self.device
|
||||
|
||||
def decode_latents(self, latents):
|
||||
if self.latent_space:
|
||||
latents = latents / self.vae.config.scaling_factor + self.vae.config.shift_factor
|
||||
latents = rearrange(latents, "b c f h w -> (b f) c h w")
|
||||
decoded_latents = self.vae.decode(latents).sample
|
||||
else:
|
||||
decoded_latents = rearrange(latents, "b c f h w -> (b f) c h w")
|
||||
return decoded_latents
|
||||
|
||||
def prepare_extra_step_kwargs(self, generator, eta):
|
||||
# prepare extra kwargs for the scheduler step, since not all schedulers have the same signature
|
||||
# eta (η) is only used with the DDIMScheduler, it will be ignored for other schedulers.
|
||||
# eta corresponds to η in DDIM paper: https://arxiv.org/abs/2010.02502
|
||||
# and should be between [0, 1]
|
||||
|
||||
accepts_eta = "eta" in set(inspect.signature(self.scheduler.step).parameters.keys())
|
||||
extra_step_kwargs = {}
|
||||
if accepts_eta:
|
||||
extra_step_kwargs["eta"] = eta
|
||||
|
||||
# check if the scheduler accepts generator
|
||||
accepts_generator = "generator" in set(inspect.signature(self.scheduler.step).parameters.keys())
|
||||
if accepts_generator:
|
||||
extra_step_kwargs["generator"] = generator
|
||||
return extra_step_kwargs
|
||||
|
||||
def check_inputs(self, height, width, callback_steps):
|
||||
assert height == width, "Height and width must be equal"
|
||||
|
||||
if height % 8 != 0 or width % 8 != 0:
|
||||
raise ValueError(f"`height` and `width` have to be divisible by 8 but are {height} and {width}.")
|
||||
|
||||
if (callback_steps is None) or (
|
||||
callback_steps is not None and (not isinstance(callback_steps, int) or callback_steps <= 0)
|
||||
):
|
||||
raise ValueError(
|
||||
f"`callback_steps` has to be a positive integer but is {callback_steps} of type"
|
||||
f" {type(callback_steps)}."
|
||||
)
|
||||
|
||||
def prepare_latents(self, batch_size, num_frames, num_channels_latents, height, width, dtype, device, generator):
|
||||
if self.latent_space:
|
||||
shape = (
|
||||
batch_size,
|
||||
num_channels_latents,
|
||||
1,
|
||||
height // self.vae_scale_factor,
|
||||
width // self.vae_scale_factor,
|
||||
)
|
||||
else:
|
||||
shape = (
|
||||
batch_size,
|
||||
num_channels_latents,
|
||||
1,
|
||||
height,
|
||||
width,
|
||||
)
|
||||
rand_device = "cpu" if device.type == "mps" else device
|
||||
latents = torch.randn(shape, generator=generator, device=rand_device, dtype=dtype).to(device)
|
||||
latents = latents.repeat(1, 1, num_frames, 1, 1)
|
||||
|
||||
# scale the initial noise by the standard deviation required by the scheduler
|
||||
latents = latents * self.scheduler.init_noise_sigma
|
||||
return latents
|
||||
|
||||
def prepare_mask_latents(
|
||||
self, mask, masked_image, height, width, dtype, device, generator, do_classifier_free_guidance
|
||||
):
|
||||
if self.latent_space:
|
||||
# resize the mask to latents shape as we concatenate the mask to the latents
|
||||
# we do that before converting to dtype to avoid breaking in case we're using cpu_offload
|
||||
# and half precision
|
||||
mask = torch.nn.functional.interpolate(
|
||||
mask, size=(height // self.vae_scale_factor, width // self.vae_scale_factor)
|
||||
)
|
||||
masked_image = masked_image.to(device=device, dtype=dtype)
|
||||
|
||||
# encode the mask image into latents space so we can concatenate it to the latents
|
||||
masked_image_latents = self.vae.encode(masked_image).latent_dist.sample(generator=generator)
|
||||
masked_image_latents = (
|
||||
masked_image_latents - self.vae.config.shift_factor
|
||||
) * self.vae.config.scaling_factor
|
||||
else:
|
||||
masked_image_latents = masked_image.to(device=device, dtype=dtype)
|
||||
|
||||
# aligning device to prevent device errors when concating it with the latent model input
|
||||
masked_image_latents = masked_image_latents.to(device=device, dtype=dtype)
|
||||
mask = mask.to(device=device, dtype=dtype)
|
||||
|
||||
# assume batch size = 1
|
||||
mask = rearrange(mask, "f c h w -> 1 c f h w")
|
||||
masked_image_latents = rearrange(masked_image_latents, "f c h w -> 1 c f h w")
|
||||
|
||||
mask = torch.cat([mask] * 2) if do_classifier_free_guidance else mask
|
||||
masked_image_latents = (
|
||||
torch.cat([masked_image_latents] * 2) if do_classifier_free_guidance else masked_image_latents
|
||||
)
|
||||
return mask, masked_image_latents
|
||||
|
||||
def prepare_image_latents(self, images, device, dtype, generator, do_classifier_free_guidance):
|
||||
images = images.to(device=device, dtype=dtype)
|
||||
if self.latent_space:
|
||||
image_latents = self.vae.encode(images).latent_dist.sample(generator=generator)
|
||||
image_latents = (image_latents - self.vae.config.shift_factor) * self.vae.config.scaling_factor
|
||||
else:
|
||||
image_latents = images
|
||||
image_latents = rearrange(image_latents, "f c h w -> 1 c f h w")
|
||||
image_latents = torch.cat([image_latents] * 2) if do_classifier_free_guidance else image_latents
|
||||
|
||||
return image_latents
|
||||
|
||||
def set_progress_bar_config(self, **kwargs):
|
||||
if not hasattr(self, "_progress_bar_config"):
|
||||
self._progress_bar_config = {}
|
||||
self._progress_bar_config.update(kwargs)
|
||||
|
||||
@staticmethod
|
||||
def recover_original_pixel_values(decoded_latents, pixel_values, masks, device, weight_dtype):
|
||||
# Combine the pixel values
|
||||
pixel_values = pixel_values.to(device=device, dtype=weight_dtype)
|
||||
masks = masks.to(device=device, dtype=weight_dtype)
|
||||
combined_pixel_values = decoded_latents * masks + pixel_values * (1 - masks)
|
||||
return combined_pixel_values
|
||||
|
||||
@staticmethod
|
||||
def pixel_values_to_images(pixel_values: torch.Tensor):
|
||||
pixel_values = rearrange(pixel_values, "f c h w -> f h w c")
|
||||
pixel_values = (pixel_values / 2 + 0.5).clamp(0, 1)
|
||||
images = (pixel_values * 255).to(torch.uint8)
|
||||
images = images.cpu().numpy()
|
||||
return images
|
||||
|
||||
def crop_audio_window(self, original_mel, start_index):
|
||||
start_idx = int(80.0 * (start_index / float(self.video_fps)))
|
||||
end_idx = start_idx + self.mel_window_length
|
||||
return original_mel[:, start_idx:end_idx].unsqueeze(0)
|
||||
|
||||
def affine_transform_video(self, video_path):
|
||||
video_frames = read_video(video_path, use_decord=False)
|
||||
faces = []
|
||||
boxes = []
|
||||
affine_matrices = []
|
||||
print(f"Affine transforming {len(video_frames)} faces...")
|
||||
for frame in tqdm.tqdm(video_frames):
|
||||
face, box, affine_matrix = self.image_processor.affine_transform(frame)
|
||||
faces.append(face)
|
||||
boxes.append(box)
|
||||
affine_matrices.append(affine_matrix)
|
||||
|
||||
faces = torch.stack(faces)
|
||||
return faces, video_frames, boxes, affine_matrices
|
||||
|
||||
def restore_video(self, faces, video_frames, boxes, affine_matrices):
|
||||
video_frames = video_frames[: faces.shape[0]]
|
||||
out_frames = []
|
||||
for index, face in enumerate(faces):
|
||||
x1, y1, x2, y2 = boxes[index]
|
||||
height = int(y2 - y1)
|
||||
width = int(x2 - x1)
|
||||
face = torchvision.transforms.functional.resize(face, size=(height, width), antialias=True)
|
||||
face = rearrange(face, "c h w -> h w c")
|
||||
face = (face / 2 + 0.5).clamp(0, 1)
|
||||
face = (face * 255).to(torch.uint8).cpu().numpy()
|
||||
out_frame = self.image_processor.restorer.restore_img(video_frames[index], face, affine_matrices[index])
|
||||
out_frames.append(out_frame)
|
||||
return np.stack(out_frames, axis=0)
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(
|
||||
self,
|
||||
video_path: str,
|
||||
audio_path: str,
|
||||
video_out_path: str,
|
||||
video_mask_path: str = None,
|
||||
num_frames: int = 16,
|
||||
video_fps: int = 25,
|
||||
audio_sample_rate: int = 16000,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
num_inference_steps: int = 50,
|
||||
guidance_scale: float = 7.5,
|
||||
weight_dtype: Optional[torch.dtype] = torch.float16,
|
||||
eta: float = 0.0,
|
||||
mask: str = "fix_mask",
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
callback: Optional[Callable[[int, int, torch.FloatTensor], None]] = None,
|
||||
callback_steps: Optional[int] = 1,
|
||||
**kwargs,
|
||||
):
|
||||
is_train = self.unet.training
|
||||
self.unet.eval()
|
||||
|
||||
# 0. Define call parameters
|
||||
batch_size = 1
|
||||
device = self._execution_device
|
||||
self.image_processor = ImageProcessor(height, mask=mask, device="cuda")
|
||||
self.set_progress_bar_config(desc=f"Sample frames: {num_frames}")
|
||||
|
||||
video_frames, original_video_frames, boxes, affine_matrices = self.affine_transform_video(video_path)
|
||||
audio_samples = read_audio(audio_path)
|
||||
|
||||
# 1. Default height and width to unet
|
||||
if self.latent_space:
|
||||
height = height or self.unet.config.sample_size * self.vae_scale_factor
|
||||
width = width or self.unet.config.sample_size * self.vae_scale_factor
|
||||
|
||||
# 2. Check inputs
|
||||
self.check_inputs(height, width, callback_steps)
|
||||
|
||||
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
|
||||
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
|
||||
# corresponds to doing no classifier free guidance.
|
||||
do_classifier_free_guidance = guidance_scale > 1.0
|
||||
|
||||
# 3. set timesteps
|
||||
self.scheduler.set_timesteps(num_inference_steps, device=device)
|
||||
timesteps = self.scheduler.timesteps
|
||||
|
||||
# 4. Prepare extra step kwargs.
|
||||
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
|
||||
|
||||
self.video_fps = video_fps
|
||||
|
||||
if self.unet.add_audio_layer:
|
||||
whisper_feature = self.audio_encoder.audio2feat(audio_path)
|
||||
whisper_chunks = self.audio_encoder.feature2chunks(feature_array=whisper_feature, fps=video_fps)
|
||||
|
||||
num_inferences = min(len(video_frames), len(whisper_chunks)) // num_frames
|
||||
else:
|
||||
num_inferences = len(video_frames) // num_frames
|
||||
|
||||
synced_video_frames = []
|
||||
masked_video_frames = []
|
||||
|
||||
save_affine_faces = False
|
||||
if save_affine_faces:
|
||||
pixel_values_faces = []
|
||||
masked_pixel_values_faces = []
|
||||
|
||||
# Prepare latent variables
|
||||
if self.latent_space:
|
||||
num_channels_latents = self.vae.config.latent_channels
|
||||
else:
|
||||
num_channels_latents = 3
|
||||
|
||||
all_latents = self.prepare_latents(
|
||||
batch_size,
|
||||
num_frames * num_inferences,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
weight_dtype,
|
||||
device,
|
||||
generator,
|
||||
)
|
||||
|
||||
for i in tqdm.tqdm(range(num_inferences), desc="Doing inference..."):
|
||||
if self.unet.add_audio_layer:
|
||||
# mel_overlap = torch.stack(mel_overlap_list[i * num_frames : (i + 1) * num_frames])
|
||||
# mel_overlap = mel_overlap.unsqueeze(0).to(device, dtype=weight_dtype)
|
||||
mel_overlap = torch.stack(whisper_chunks[i * num_frames : (i + 1) * num_frames])
|
||||
mel_overlap = mel_overlap.to(device, dtype=weight_dtype)
|
||||
if do_classifier_free_guidance:
|
||||
empty_mel_overlap = torch.zeros_like(mel_overlap)
|
||||
mel_overlap = torch.cat([empty_mel_overlap, mel_overlap])
|
||||
else:
|
||||
mel_overlap = None
|
||||
inference_video_frames = video_frames[i * num_frames : (i + 1) * num_frames]
|
||||
latents = all_latents[:, :, i * num_frames : (i + 1) * num_frames]
|
||||
pixel_values, masked_pixel_values, masks = self.image_processor.prepare_masks_and_masked_images(
|
||||
inference_video_frames, affine_transform=False
|
||||
)
|
||||
|
||||
if save_affine_faces:
|
||||
pixel_values_faces.append(pixel_values)
|
||||
masked_pixel_values_faces.append(masked_pixel_values)
|
||||
|
||||
# 7. Prepare mask latent variables
|
||||
mask_latents, masked_image_latents = self.prepare_mask_latents(
|
||||
masks,
|
||||
masked_pixel_values,
|
||||
height,
|
||||
width,
|
||||
weight_dtype,
|
||||
device,
|
||||
generator,
|
||||
do_classifier_free_guidance,
|
||||
)
|
||||
|
||||
# 8. Prepare image latents
|
||||
image_latents = self.prepare_image_latents(
|
||||
pixel_values,
|
||||
device,
|
||||
weight_dtype,
|
||||
generator,
|
||||
do_classifier_free_guidance,
|
||||
)
|
||||
|
||||
# 9. Denoising loop
|
||||
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for j, t in enumerate(timesteps):
|
||||
# expand the latents if we are doing classifier free guidance
|
||||
latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents
|
||||
|
||||
# concat latents, mask, masked_image_latents in the channel dimension
|
||||
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
|
||||
latent_model_input = torch.cat(
|
||||
[latent_model_input, mask_latents, masked_image_latents, image_latents], dim=1
|
||||
)
|
||||
|
||||
# predict the noise residual
|
||||
noise_pred = self.unet(latent_model_input, t, encoder_hidden_states=mel_overlap).sample
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs).prev_sample
|
||||
|
||||
# call the callback, if provided
|
||||
if j == len(timesteps) - 1 or (
|
||||
(j + 1) > num_warmup_steps and (j + 1) % self.scheduler.order == 0
|
||||
):
|
||||
progress_bar.update()
|
||||
if callback is not None and j % callback_steps == 0:
|
||||
callback(j, t, latents)
|
||||
|
||||
# Recover the pixel values
|
||||
decoded_latents = self.decode_latents(latents)
|
||||
decoded_latents = self.recover_original_pixel_values(
|
||||
decoded_latents, pixel_values, 1 - masks, device, weight_dtype
|
||||
)
|
||||
synced_video_frames.append(decoded_latents)
|
||||
masked_video_frames.append(masked_pixel_values)
|
||||
|
||||
synced_video_frames = self.restore_video(
|
||||
torch.cat(synced_video_frames), original_video_frames, boxes, affine_matrices
|
||||
)
|
||||
masked_video_frames = self.restore_video(
|
||||
torch.cat(masked_video_frames), original_video_frames, boxes, affine_matrices
|
||||
)
|
||||
|
||||
audio_samples_remain_length = int(synced_video_frames.shape[0] / video_fps * audio_sample_rate)
|
||||
audio_samples = audio_samples[:audio_samples_remain_length].cpu().numpy()
|
||||
|
||||
if is_train:
|
||||
self.unet.train()
|
||||
|
||||
temp_dir = "temp"
|
||||
if os.path.exists(temp_dir):
|
||||
shutil.rmtree(temp_dir)
|
||||
os.makedirs(temp_dir, exist_ok=True)
|
||||
|
||||
if save_affine_faces:
|
||||
pixel_values_faces = torch.cat(pixel_values_faces)
|
||||
masked_pixel_values_faces = torch.cat(masked_pixel_values_faces)
|
||||
|
||||
pixel_values_faces = self.pixel_values_to_images(pixel_values_faces)
|
||||
masked_pixel_values_faces = self.pixel_values_to_images(masked_pixel_values_faces)
|
||||
|
||||
write_video("affine_faces.mp4", pixel_values_faces, fps=25)
|
||||
write_video("masked_affine_faces.mp4", masked_pixel_values_faces, fps=25)
|
||||
|
||||
write_video(os.path.join(temp_dir, "video.mp4"), synced_video_frames, fps=25)
|
||||
# write_video(video_mask_path, masked_video_frames, fps=25)
|
||||
|
||||
sf.write(os.path.join(temp_dir, "audio.wav"), audio_samples, audio_sample_rate)
|
||||
|
||||
command = f"ffmpeg -y -loglevel error -nostdin -i {os.path.join(temp_dir, 'video.mp4')} -i {os.path.join(temp_dir, 'audio.wav')} -c:v libx264 -c:a aac -q:v 0 -q:a 0 {video_out_path}"
|
||||
subprocess.run(command, shell=True)
|
||||
# Adapted from https://github.com/guoyww/AnimateDiff/blob/main/animatediff/pipelines/pipeline_animation.py
|
||||
|
||||
import inspect
|
||||
import os
|
||||
import shutil
|
||||
from typing import Callable, List, Optional, Union
|
||||
import subprocess
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
|
||||
from packaging import version
|
||||
|
||||
from diffusers.configuration_utils import FrozenDict
|
||||
from diffusers.models import AutoencoderKL
|
||||
from diffusers.pipelines import DiffusionPipeline
|
||||
from diffusers.schedulers import (
|
||||
DDIMScheduler,
|
||||
DPMSolverMultistepScheduler,
|
||||
EulerAncestralDiscreteScheduler,
|
||||
EulerDiscreteScheduler,
|
||||
LMSDiscreteScheduler,
|
||||
PNDMScheduler,
|
||||
)
|
||||
from diffusers.utils import deprecate, logging
|
||||
|
||||
from einops import rearrange
|
||||
import cv2
|
||||
|
||||
from ..models.unet import UNet3DConditionModel
|
||||
from ..utils.util import read_video, read_audio, write_video, check_ffmpeg_installed
|
||||
from ..utils.image_processor import ImageProcessor, load_fixed_mask
|
||||
from ..whisper.audio2feature import Audio2Feature
|
||||
import tqdm
|
||||
import soundfile as sf
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
class LipsyncPipeline(DiffusionPipeline):
|
||||
_optional_components = []
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vae: AutoencoderKL,
|
||||
audio_encoder: Audio2Feature,
|
||||
denoising_unet: UNet3DConditionModel,
|
||||
scheduler: Union[
|
||||
DDIMScheduler,
|
||||
PNDMScheduler,
|
||||
LMSDiscreteScheduler,
|
||||
EulerDiscreteScheduler,
|
||||
EulerAncestralDiscreteScheduler,
|
||||
DPMSolverMultistepScheduler,
|
||||
],
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
if hasattr(scheduler.config, "steps_offset") and scheduler.config.steps_offset != 1:
|
||||
deprecation_message = (
|
||||
f"The configuration file of this scheduler: {scheduler} is outdated. `steps_offset`"
|
||||
f" should be set to 1 instead of {scheduler.config.steps_offset}. Please make sure "
|
||||
"to update the config accordingly as leaving `steps_offset` might led to incorrect results"
|
||||
" in future versions. If you have downloaded this checkpoint from the Hugging Face Hub,"
|
||||
" it would be very nice if you could open a Pull request for the `scheduler/scheduler_config.json`"
|
||||
" file"
|
||||
)
|
||||
deprecate("steps_offset!=1", "1.0.0", deprecation_message, standard_warn=False)
|
||||
new_config = dict(scheduler.config)
|
||||
new_config["steps_offset"] = 1
|
||||
scheduler._internal_dict = FrozenDict(new_config)
|
||||
|
||||
if hasattr(scheduler.config, "clip_sample") and scheduler.config.clip_sample is True:
|
||||
deprecation_message = (
|
||||
f"The configuration file of this scheduler: {scheduler} has not set the configuration `clip_sample`."
|
||||
" `clip_sample` should be set to False in the configuration file. Please make sure to update the"
|
||||
" config accordingly as not setting `clip_sample` in the config might lead to incorrect results in"
|
||||
" future versions. If you have downloaded this checkpoint from the Hugging Face Hub, it would be very"
|
||||
" nice if you could open a Pull request for the `scheduler/scheduler_config.json` file"
|
||||
)
|
||||
deprecate("clip_sample not set", "1.0.0", deprecation_message, standard_warn=False)
|
||||
new_config = dict(scheduler.config)
|
||||
new_config["clip_sample"] = False
|
||||
scheduler._internal_dict = FrozenDict(new_config)
|
||||
|
||||
is_unet_version_less_0_9_0 = hasattr(denoising_unet.config, "_diffusers_version") and version.parse(
|
||||
version.parse(denoising_unet.config._diffusers_version).base_version
|
||||
) < version.parse("0.9.0.dev0")
|
||||
is_unet_sample_size_less_64 = (
|
||||
hasattr(denoising_unet.config, "sample_size") and denoising_unet.config.sample_size < 64
|
||||
)
|
||||
if is_unet_version_less_0_9_0 and is_unet_sample_size_less_64:
|
||||
deprecation_message = (
|
||||
"The configuration file of the unet has set the default `sample_size` to smaller than"
|
||||
" 64 which seems highly unlikely. If your checkpoint is a fine-tuned version of any of the"
|
||||
" following: \n- CompVis/stable-diffusion-v1-4 \n- CompVis/stable-diffusion-v1-3 \n-"
|
||||
" CompVis/stable-diffusion-v1-2 \n- CompVis/stable-diffusion-v1-1 \n- runwayml/stable-diffusion-v1-5"
|
||||
" \n- runwayml/stable-diffusion-inpainting \n you should change 'sample_size' to 64 in the"
|
||||
" configuration file. Please make sure to update the config accordingly as leaving `sample_size=32`"
|
||||
" in the config might lead to incorrect results in future versions. If you have downloaded this"
|
||||
" checkpoint from the Hugging Face Hub, it would be very nice if you could open a Pull request for"
|
||||
" the `unet/config.json` file"
|
||||
)
|
||||
deprecate("sample_size<64", "1.0.0", deprecation_message, standard_warn=False)
|
||||
new_config = dict(denoising_unet.config)
|
||||
new_config["sample_size"] = 64
|
||||
denoising_unet._internal_dict = FrozenDict(new_config)
|
||||
|
||||
self.register_modules(
|
||||
vae=vae,
|
||||
audio_encoder=audio_encoder,
|
||||
denoising_unet=denoising_unet,
|
||||
scheduler=scheduler,
|
||||
)
|
||||
|
||||
self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
|
||||
|
||||
self.set_progress_bar_config(desc="Steps")
|
||||
|
||||
def enable_vae_slicing(self):
|
||||
self.vae.enable_slicing()
|
||||
|
||||
def disable_vae_slicing(self):
|
||||
self.vae.disable_slicing()
|
||||
|
||||
@property
|
||||
def _execution_device(self):
|
||||
if self.device != torch.device("meta") or not hasattr(self.denoising_unet, "_hf_hook"):
|
||||
return self.device
|
||||
for module in self.denoising_unet.modules():
|
||||
if (
|
||||
hasattr(module, "_hf_hook")
|
||||
and hasattr(module._hf_hook, "execution_device")
|
||||
and module._hf_hook.execution_device is not None
|
||||
):
|
||||
return torch.device(module._hf_hook.execution_device)
|
||||
return self.device
|
||||
|
||||
def decode_latents(self, latents):
|
||||
latents = latents / self.vae.config.scaling_factor + self.vae.config.shift_factor
|
||||
latents = rearrange(latents, "b c f h w -> (b f) c h w")
|
||||
decoded_latents = self.vae.decode(latents).sample
|
||||
return decoded_latents
|
||||
|
||||
def prepare_extra_step_kwargs(self, generator, eta):
|
||||
# prepare extra kwargs for the scheduler step, since not all schedulers have the same signature
|
||||
# eta (η) is only used with the DDIMScheduler, it will be ignored for other schedulers.
|
||||
# eta corresponds to η in DDIM paper: https://arxiv.org/abs/2010.02502
|
||||
# and should be between [0, 1]
|
||||
|
||||
accepts_eta = "eta" in set(inspect.signature(self.scheduler.step).parameters.keys())
|
||||
extra_step_kwargs = {}
|
||||
if accepts_eta:
|
||||
extra_step_kwargs["eta"] = eta
|
||||
|
||||
# check if the scheduler accepts generator
|
||||
accepts_generator = "generator" in set(inspect.signature(self.scheduler.step).parameters.keys())
|
||||
if accepts_generator:
|
||||
extra_step_kwargs["generator"] = generator
|
||||
return extra_step_kwargs
|
||||
|
||||
def check_inputs(self, height, width, callback_steps):
|
||||
assert height == width, "Height and width must be equal"
|
||||
|
||||
if height % 8 != 0 or width % 8 != 0:
|
||||
raise ValueError(f"`height` and `width` have to be divisible by 8 but are {height} and {width}.")
|
||||
|
||||
if (callback_steps is None) or (
|
||||
callback_steps is not None and (not isinstance(callback_steps, int) or callback_steps <= 0)
|
||||
):
|
||||
raise ValueError(
|
||||
f"`callback_steps` has to be a positive integer but is {callback_steps} of type"
|
||||
f" {type(callback_steps)}."
|
||||
)
|
||||
|
||||
def prepare_latents(self, batch_size, num_frames, num_channels_latents, height, width, dtype, device, generator):
|
||||
shape = (
|
||||
batch_size,
|
||||
num_channels_latents,
|
||||
1,
|
||||
height // self.vae_scale_factor,
|
||||
width // self.vae_scale_factor,
|
||||
)
|
||||
rand_device = "cpu" if device.type == "mps" else device
|
||||
latents = torch.randn(shape, generator=generator, device=rand_device, dtype=dtype).to(device)
|
||||
latents = latents.repeat(1, 1, num_frames, 1, 1)
|
||||
|
||||
# scale the initial noise by the standard deviation required by the scheduler
|
||||
latents = latents * self.scheduler.init_noise_sigma
|
||||
return latents
|
||||
|
||||
def prepare_mask_latents(
|
||||
self, mask, masked_image, height, width, dtype, device, generator, do_classifier_free_guidance
|
||||
):
|
||||
# resize the mask to latents shape as we concatenate the mask to the latents
|
||||
# we do that before converting to dtype to avoid breaking in case we're using cpu_offload
|
||||
# and half precision
|
||||
mask = torch.nn.functional.interpolate(
|
||||
mask, size=(height // self.vae_scale_factor, width // self.vae_scale_factor)
|
||||
)
|
||||
masked_image = masked_image.to(device=device, dtype=dtype)
|
||||
|
||||
# encode the mask image into latents space so we can concatenate it to the latents
|
||||
masked_image_latents = self.vae.encode(masked_image).latent_dist.sample(generator=generator)
|
||||
masked_image_latents = (masked_image_latents - self.vae.config.shift_factor) * self.vae.config.scaling_factor
|
||||
|
||||
# aligning device to prevent device errors when concating it with the latent model input
|
||||
masked_image_latents = masked_image_latents.to(device=device, dtype=dtype)
|
||||
mask = mask.to(device=device, dtype=dtype)
|
||||
|
||||
# assume batch size = 1
|
||||
mask = rearrange(mask, "f c h w -> 1 c f h w")
|
||||
masked_image_latents = rearrange(masked_image_latents, "f c h w -> 1 c f h w")
|
||||
|
||||
mask = torch.cat([mask] * 2) if do_classifier_free_guidance else mask
|
||||
masked_image_latents = (
|
||||
torch.cat([masked_image_latents] * 2) if do_classifier_free_guidance else masked_image_latents
|
||||
)
|
||||
return mask, masked_image_latents
|
||||
|
||||
def prepare_image_latents(self, images, device, dtype, generator, do_classifier_free_guidance):
|
||||
images = images.to(device=device, dtype=dtype)
|
||||
image_latents = self.vae.encode(images).latent_dist.sample(generator=generator)
|
||||
image_latents = (image_latents - self.vae.config.shift_factor) * self.vae.config.scaling_factor
|
||||
image_latents = rearrange(image_latents, "f c h w -> 1 c f h w")
|
||||
image_latents = torch.cat([image_latents] * 2) if do_classifier_free_guidance else image_latents
|
||||
|
||||
return image_latents
|
||||
|
||||
def set_progress_bar_config(self, **kwargs):
|
||||
if not hasattr(self, "_progress_bar_config"):
|
||||
self._progress_bar_config = {}
|
||||
self._progress_bar_config.update(kwargs)
|
||||
|
||||
@staticmethod
|
||||
def paste_surrounding_pixels_back(decoded_latents, pixel_values, masks, device, weight_dtype):
|
||||
# Paste the surrounding pixels back, because we only want to change the mouth region
|
||||
pixel_values = pixel_values.to(device=device, dtype=weight_dtype)
|
||||
masks = masks.to(device=device, dtype=weight_dtype)
|
||||
combined_pixel_values = decoded_latents * masks + pixel_values * (1 - masks)
|
||||
return combined_pixel_values
|
||||
|
||||
@staticmethod
|
||||
def pixel_values_to_images(pixel_values: torch.Tensor):
|
||||
pixel_values = rearrange(pixel_values, "f c h w -> f h w c")
|
||||
pixel_values = (pixel_values / 2 + 0.5).clamp(0, 1)
|
||||
images = (pixel_values * 255).to(torch.uint8)
|
||||
images = images.cpu().numpy()
|
||||
return images
|
||||
|
||||
def affine_transform_video(self, video_frames: np.ndarray):
|
||||
faces = []
|
||||
boxes = []
|
||||
affine_matrices = []
|
||||
print(f"Affine transforming {len(video_frames)} faces...")
|
||||
for frame in tqdm.tqdm(video_frames):
|
||||
face, box, affine_matrix = self.image_processor.affine_transform(frame)
|
||||
faces.append(face)
|
||||
boxes.append(box)
|
||||
affine_matrices.append(affine_matrix)
|
||||
|
||||
faces = torch.stack(faces)
|
||||
return faces, boxes, affine_matrices
|
||||
|
||||
def restore_video(self, faces, video_frames, boxes, affine_matrices):
|
||||
video_frames = video_frames[: faces.shape[0]]
|
||||
out_frames = []
|
||||
print(f"Restoring {len(faces)} faces...")
|
||||
for index, face in enumerate(tqdm.tqdm(faces)):
|
||||
x1, y1, x2, y2 = boxes[index]
|
||||
height = int(y2 - y1)
|
||||
width = int(x2 - x1)
|
||||
face = torchvision.transforms.functional.resize(face, size=(height, width), antialias=True)
|
||||
face = rearrange(face, "c h w -> h w c")
|
||||
face = (face / 2 + 0.5).clamp(0, 1)
|
||||
face = (face * 255).to(torch.uint8).cpu().numpy()
|
||||
# face = cv2.resize(face, (width, height), interpolation=cv2.INTER_LANCZOS4)
|
||||
out_frame = self.image_processor.restorer.restore_img(video_frames[index], face, affine_matrices[index])
|
||||
out_frames.append(out_frame)
|
||||
return np.stack(out_frames, axis=0)
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(
|
||||
self,
|
||||
video_path: str,
|
||||
audio_path: str,
|
||||
video_out_path: str,
|
||||
video_mask_path: str = None,
|
||||
num_frames: int = 16,
|
||||
video_fps: int = 25,
|
||||
audio_sample_rate: int = 16000,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
num_inference_steps: int = 20,
|
||||
guidance_scale: float = 1.5,
|
||||
weight_dtype: Optional[torch.dtype] = torch.float16,
|
||||
eta: float = 0.0,
|
||||
mask: str = "fix_mask",
|
||||
mask_image_path: str = "latentsync/utils/mask.png",
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
callback: Optional[Callable[[int, int, torch.FloatTensor], None]] = None,
|
||||
callback_steps: Optional[int] = 1,
|
||||
**kwargs,
|
||||
):
|
||||
is_train = self.denoising_unet.training
|
||||
self.denoising_unet.eval()
|
||||
|
||||
check_ffmpeg_installed()
|
||||
|
||||
# 0. Define call parameters
|
||||
batch_size = 1
|
||||
device = self._execution_device
|
||||
mask_image = load_fixed_mask(height, mask_image_path)
|
||||
self.image_processor = ImageProcessor(height, mask=mask, device="cuda", mask_image=mask_image)
|
||||
self.set_progress_bar_config(desc=f"Sample frames: {num_frames}")
|
||||
|
||||
# 1. Default height and width to unet
|
||||
height = height or self.denoising_unet.config.sample_size * self.vae_scale_factor
|
||||
width = width or self.denoising_unet.config.sample_size * self.vae_scale_factor
|
||||
|
||||
# 2. Check inputs
|
||||
self.check_inputs(height, width, callback_steps)
|
||||
|
||||
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
|
||||
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
|
||||
# corresponds to doing no classifier free guidance.
|
||||
do_classifier_free_guidance = guidance_scale > 1.0
|
||||
|
||||
# 3. set timesteps
|
||||
self.scheduler.set_timesteps(num_inference_steps, device=device)
|
||||
timesteps = self.scheduler.timesteps
|
||||
|
||||
# 4. Prepare extra step kwargs.
|
||||
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
|
||||
|
||||
whisper_feature = self.audio_encoder.audio2feat(audio_path)
|
||||
whisper_chunks = self.audio_encoder.feature2chunks(feature_array=whisper_feature, fps=video_fps)
|
||||
|
||||
audio_samples = read_audio(audio_path)
|
||||
video_frames = read_video(video_path, use_decord=False)
|
||||
|
||||
num_inferences = min(len(video_frames), len(whisper_chunks)) // num_frames
|
||||
video_frames = video_frames[: num_inferences * num_frames]
|
||||
faces, boxes, affine_matrices = self.affine_transform_video(video_frames)
|
||||
|
||||
synced_video_frames = []
|
||||
masked_video_frames = []
|
||||
|
||||
num_channels_latents = self.vae.config.latent_channels
|
||||
|
||||
# Prepare latent variables
|
||||
all_latents = self.prepare_latents(
|
||||
batch_size,
|
||||
num_frames * num_inferences,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
weight_dtype,
|
||||
device,
|
||||
generator,
|
||||
)
|
||||
|
||||
for i in tqdm.tqdm(range(num_inferences), desc="Doing inference..."):
|
||||
if self.denoising_unet.add_audio_layer:
|
||||
audio_embeds = torch.stack(whisper_chunks[i * num_frames : (i + 1) * num_frames])
|
||||
audio_embeds = audio_embeds.to(device, dtype=weight_dtype)
|
||||
if do_classifier_free_guidance:
|
||||
null_audio_embeds = torch.zeros_like(audio_embeds)
|
||||
audio_embeds = torch.cat([null_audio_embeds, audio_embeds])
|
||||
else:
|
||||
audio_embeds = None
|
||||
inference_faces = faces[i * num_frames : (i + 1) * num_frames]
|
||||
latents = all_latents[:, :, i * num_frames : (i + 1) * num_frames]
|
||||
ref_pixel_values, masked_pixel_values, masks = self.image_processor.prepare_masks_and_masked_images(
|
||||
inference_faces, affine_transform=False
|
||||
)
|
||||
|
||||
# 7. Prepare mask latent variables
|
||||
mask_latents, masked_image_latents = self.prepare_mask_latents(
|
||||
masks,
|
||||
masked_pixel_values,
|
||||
height,
|
||||
width,
|
||||
weight_dtype,
|
||||
device,
|
||||
generator,
|
||||
do_classifier_free_guidance,
|
||||
)
|
||||
|
||||
# 8. Prepare image latents
|
||||
ref_latents = self.prepare_image_latents(
|
||||
ref_pixel_values,
|
||||
device,
|
||||
weight_dtype,
|
||||
generator,
|
||||
do_classifier_free_guidance,
|
||||
)
|
||||
|
||||
# 9. Denoising loop
|
||||
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for j, t in enumerate(timesteps):
|
||||
# expand the latents if we are doing classifier free guidance
|
||||
denoising_unet_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents
|
||||
|
||||
denoising_unet_input = self.scheduler.scale_model_input(denoising_unet_input, t)
|
||||
|
||||
# concat latents, mask, masked_image_latents in the channel dimension
|
||||
denoising_unet_input = torch.cat(
|
||||
[denoising_unet_input, mask_latents, masked_image_latents, ref_latents], dim=1
|
||||
)
|
||||
|
||||
# predict the noise residual
|
||||
noise_pred = self.denoising_unet(
|
||||
denoising_unet_input, t, encoder_hidden_states=audio_embeds
|
||||
).sample
|
||||
|
||||
# perform guidance
|
||||
if do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_audio = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_audio - noise_pred_uncond)
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs).prev_sample
|
||||
|
||||
# call the callback, if provided
|
||||
if j == len(timesteps) - 1 or ((j + 1) > num_warmup_steps and (j + 1) % self.scheduler.order == 0):
|
||||
progress_bar.update()
|
||||
if callback is not None and j % callback_steps == 0:
|
||||
callback(j, t, latents)
|
||||
|
||||
# Recover the pixel values
|
||||
decoded_latents = self.decode_latents(latents)
|
||||
decoded_latents = self.paste_surrounding_pixels_back(
|
||||
decoded_latents, ref_pixel_values, 1 - masks, device, weight_dtype
|
||||
)
|
||||
synced_video_frames.append(decoded_latents)
|
||||
# masked_video_frames.append(masked_pixel_values)
|
||||
|
||||
synced_video_frames = self.restore_video(
|
||||
torch.cat(synced_video_frames), video_frames, boxes, affine_matrices
|
||||
)
|
||||
# masked_video_frames = self.restore_video(
|
||||
# torch.cat(masked_video_frames), video_frames, boxes, affine_matrices
|
||||
# )
|
||||
|
||||
audio_samples_remain_length = int(synced_video_frames.shape[0] / video_fps * audio_sample_rate)
|
||||
audio_samples = audio_samples[:audio_samples_remain_length].cpu().numpy()
|
||||
|
||||
if is_train:
|
||||
self.denoising_unet.train()
|
||||
|
||||
temp_dir = "temp"
|
||||
if os.path.exists(temp_dir):
|
||||
shutil.rmtree(temp_dir)
|
||||
os.makedirs(temp_dir, exist_ok=True)
|
||||
|
||||
write_video(os.path.join(temp_dir, "video.mp4"), synced_video_frames, fps=25)
|
||||
# write_video(video_mask_path, masked_video_frames, fps=25)
|
||||
|
||||
sf.write(os.path.join(temp_dir, "audio.wav"), audio_samples, audio_sample_rate)
|
||||
|
||||
command = f"ffmpeg -y -loglevel error -nostdin -i {os.path.join(temp_dir, 'video.mp4')} -i {os.path.join(temp_dir, 'audio.wav')} -c:v libx264 -c:a aac -q:v 0 -q:a 0 {video_out_path}"
|
||||
subprocess.run(command, shell=True)
|
||||
|
||||
@@ -0,0 +1,65 @@
|
||||
# Copyright (c) 2024 Bytedance Ltd. and/or its affiliates
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
from .third_party.VideoMAEv2.utils import load_videomae_model
|
||||
|
||||
|
||||
class TREPALoss:
|
||||
def __init__(
|
||||
self,
|
||||
device="cuda",
|
||||
ckpt_path="checkpoints/auxiliary/vit_g_hybrid_pt_1200e_ssv2_ft.pth",
|
||||
with_cp=False,
|
||||
):
|
||||
self.model = load_videomae_model(device, ckpt_path, with_cp).eval().to(dtype=torch.float16)
|
||||
self.model.requires_grad_(False)
|
||||
|
||||
def __call__(self, videos_fake, videos_real):
|
||||
batch_size = videos_fake.shape[0]
|
||||
num_frames = videos_fake.shape[2]
|
||||
videos_fake = rearrange(videos_fake.clone(), "b c f h w -> (b f) c h w")
|
||||
videos_real = rearrange(videos_real.clone(), "b c f h w -> (b f) c h w")
|
||||
|
||||
videos_fake = F.interpolate(videos_fake, size=(224, 224), mode="bilinear")
|
||||
videos_real = F.interpolate(videos_real, size=(224, 224), mode="bilinear")
|
||||
|
||||
videos_fake = rearrange(videos_fake, "(b f) c h w -> b c f h w", f=num_frames)
|
||||
videos_real = rearrange(videos_real, "(b f) c h w -> b c f h w", f=num_frames)
|
||||
|
||||
# Because input pixel range is [-1, 1], and model expects pixel range to be [0, 1]
|
||||
videos_fake = (videos_fake / 2 + 0.5).clamp(0, 1)
|
||||
videos_real = (videos_real / 2 + 0.5).clamp(0, 1)
|
||||
|
||||
feats_fake = self.model.forward_features(videos_fake)
|
||||
feats_real = self.model.forward_features(videos_real)
|
||||
|
||||
feats_fake = F.normalize(feats_fake, p=2, dim=1)
|
||||
feats_real = F.normalize(feats_real, p=2, dim=1)
|
||||
|
||||
return F.mse_loss(feats_fake, feats_real)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
torch.manual_seed(42)
|
||||
|
||||
# input shape: (b, c, f, h, w)
|
||||
videos_fake = torch.randn(2, 3, 16, 256, 256, requires_grad=True).to(device="cuda", dtype=torch.float16)
|
||||
videos_real = torch.randn(2, 3, 16, 256, 256, requires_grad=True).to(device="cuda", dtype=torch.float16)
|
||||
|
||||
trepa_loss = TREPALoss(device="cuda", with_cp=True)
|
||||
loss = trepa_loss(videos_fake, videos_real)
|
||||
print(loss)
|
||||
+82
-81
@@ -1,81 +1,82 @@
|
||||
import os
|
||||
import torch
|
||||
import requests
|
||||
from tqdm import tqdm
|
||||
from torchvision import transforms
|
||||
from .videomaev2_finetune import vit_giant_patch14_224
|
||||
|
||||
def to_normalized_float_tensor(vid):
|
||||
return vid.permute(3, 0, 1, 2).to(torch.float32) / 255
|
||||
|
||||
|
||||
# NOTE: for those functions, which generally expect mini-batches, we keep them
|
||||
# as non-minibatch so that they are applied as if they were 4d (thus image).
|
||||
# this way, we only apply the transformation in the spatial domain
|
||||
def resize(vid, size, interpolation='bilinear'):
|
||||
# NOTE: using bilinear interpolation because we don't work on minibatches
|
||||
# at this level
|
||||
scale = None
|
||||
if isinstance(size, int):
|
||||
scale = float(size) / min(vid.shape[-2:])
|
||||
size = None
|
||||
return torch.nn.functional.interpolate(
|
||||
vid,
|
||||
size=size,
|
||||
scale_factor=scale,
|
||||
mode=interpolation,
|
||||
align_corners=False)
|
||||
|
||||
|
||||
class ToFloatTensorInZeroOne(object):
|
||||
def __call__(self, vid):
|
||||
return to_normalized_float_tensor(vid)
|
||||
|
||||
|
||||
class Resize(object):
|
||||
def __init__(self, size):
|
||||
self.size = size
|
||||
def __call__(self, vid):
|
||||
return resize(vid, self.size)
|
||||
|
||||
def preprocess_videomae(videos):
|
||||
transform = transforms.Compose(
|
||||
[ToFloatTensorInZeroOne(),
|
||||
Resize((224, 224))])
|
||||
return torch.stack([transform(f) for f in torch.from_numpy(videos)])
|
||||
|
||||
|
||||
def load_videomae_model(device, ckpt_path=None):
|
||||
if ckpt_path is None:
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
ckpt_path = os.path.join(current_dir, 'vit_g_hybrid_pt_1200e_ssv2_ft.pth')
|
||||
|
||||
if not os.path.exists(ckpt_path):
|
||||
# download the ckpt to the path
|
||||
ckpt_url = 'https://pjlab-gvm-data.oss-cn-shanghai.aliyuncs.com/internvideo/videomaev2/vit_g_hybrid_pt_1200e_ssv2_ft.pth'
|
||||
response = requests.get(ckpt_url, stream=True, allow_redirects=True)
|
||||
total_size = int(response.headers.get("content-length", 0))
|
||||
block_size = 1024
|
||||
|
||||
with tqdm(total=total_size, unit="B", unit_scale=True) as progress_bar:
|
||||
with open(ckpt_path, "wb") as fw:
|
||||
for data in response.iter_content(block_size):
|
||||
progress_bar.update(len(data))
|
||||
fw.write(data)
|
||||
|
||||
model = vit_giant_patch14_224(
|
||||
img_size=224,
|
||||
pretrained=False,
|
||||
num_classes=174,
|
||||
all_frames=16,
|
||||
tubelet_size=2,
|
||||
drop_path_rate=0.3,
|
||||
use_mean_pooling=True)
|
||||
|
||||
ckpt = torch.load(ckpt_path, map_location='cpu')
|
||||
for model_key in ['model', 'module']:
|
||||
if model_key in ckpt:
|
||||
ckpt = ckpt[model_key]
|
||||
break
|
||||
model.load_state_dict(ckpt)
|
||||
return model.to(device)
|
||||
import os
|
||||
import torch
|
||||
import requests
|
||||
from tqdm import tqdm
|
||||
from torchvision import transforms
|
||||
from .videomaev2_finetune import vit_giant_patch14_224
|
||||
|
||||
|
||||
def to_normalized_float_tensor(vid):
|
||||
return vid.permute(3, 0, 1, 2).to(torch.float32) / 255
|
||||
|
||||
|
||||
# NOTE: for those functions, which generally expect mini-batches, we keep them
|
||||
# as non-minibatch so that they are applied as if they were 4d (thus image).
|
||||
# this way, we only apply the transformation in the spatial domain
|
||||
def resize(vid, size, interpolation="bilinear"):
|
||||
# NOTE: using bilinear interpolation because we don't work on minibatches
|
||||
# at this level
|
||||
scale = None
|
||||
if isinstance(size, int):
|
||||
scale = float(size) / min(vid.shape[-2:])
|
||||
size = None
|
||||
return torch.nn.functional.interpolate(vid, size=size, scale_factor=scale, mode=interpolation, align_corners=False)
|
||||
|
||||
|
||||
class ToFloatTensorInZeroOne(object):
|
||||
def __call__(self, vid):
|
||||
return to_normalized_float_tensor(vid)
|
||||
|
||||
|
||||
class Resize(object):
|
||||
def __init__(self, size):
|
||||
self.size = size
|
||||
|
||||
def __call__(self, vid):
|
||||
return resize(vid, self.size)
|
||||
|
||||
|
||||
def preprocess_videomae(videos):
|
||||
transform = transforms.Compose([ToFloatTensorInZeroOne(), Resize((224, 224))])
|
||||
return torch.stack([transform(f) for f in torch.from_numpy(videos)])
|
||||
|
||||
|
||||
def load_videomae_model(device, ckpt_path=None, with_cp=False):
|
||||
if ckpt_path is None:
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
ckpt_path = os.path.join(current_dir, "vit_g_hybrid_pt_1200e_ssv2_ft.pth")
|
||||
|
||||
if not os.path.exists(ckpt_path):
|
||||
# download the ckpt to the path
|
||||
ckpt_url = "https://pjlab-gvm-data.oss-cn-shanghai.aliyuncs.com/internvideo/videomaev2/vit_g_hybrid_pt_1200e_ssv2_ft.pth"
|
||||
response = requests.get(ckpt_url, stream=True, allow_redirects=True)
|
||||
total_size = int(response.headers.get("content-length", 0))
|
||||
block_size = 1024
|
||||
|
||||
with tqdm(total=total_size, unit="B", unit_scale=True) as progress_bar:
|
||||
with open(ckpt_path, "wb") as fw:
|
||||
for data in response.iter_content(block_size):
|
||||
progress_bar.update(len(data))
|
||||
fw.write(data)
|
||||
|
||||
model = vit_giant_patch14_224(
|
||||
img_size=224,
|
||||
pretrained=False,
|
||||
num_classes=174,
|
||||
all_frames=16,
|
||||
tubelet_size=2,
|
||||
drop_path_rate=0.3,
|
||||
use_mean_pooling=True,
|
||||
with_cp=with_cp,
|
||||
)
|
||||
|
||||
ckpt = torch.load(ckpt_path, map_location="cpu", weights_only=True)
|
||||
for model_key in ["model", "module"]:
|
||||
if model_key in ckpt:
|
||||
ckpt = ckpt[model_key]
|
||||
break
|
||||
model.load_state_dict(ckpt)
|
||||
|
||||
del ckpt
|
||||
torch.cuda.empty_cache()
|
||||
return model.to(device)
|
||||
|
||||
+543
-539
File diff suppressed because it is too large
Load Diff
+469
-469
@@ -1,469 +1,469 @@
|
||||
# --------------------------------------------------------
|
||||
# Based on BEiT, timm, DINO and DeiT code bases
|
||||
# https://github.com/microsoft/unilm/tree/master/beit
|
||||
# https://github.com/rwightman/pytorch-image-models/tree/master/timm
|
||||
# https://github.com/facebookresearch/deit
|
||||
# https://github.com/facebookresearch/dino
|
||||
# --------------------------------------------------------'
|
||||
from functools import partial
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.utils.checkpoint as cp
|
||||
|
||||
from .videomaev2_finetune import (
|
||||
Block,
|
||||
PatchEmbed,
|
||||
_cfg,
|
||||
get_sinusoid_encoding_table,
|
||||
)
|
||||
|
||||
from .videomaev2_finetune import trunc_normal_ as __call_trunc_normal_
|
||||
|
||||
def trunc_normal_(tensor, mean=0., std=1.):
|
||||
__call_trunc_normal_(tensor, mean=mean, std=std, a=-std, b=std)
|
||||
|
||||
|
||||
class PretrainVisionTransformerEncoder(nn.Module):
|
||||
""" Vision Transformer with support for patch or hybrid CNN input stage
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
img_size=224,
|
||||
patch_size=16,
|
||||
in_chans=3,
|
||||
num_classes=0,
|
||||
embed_dim=768,
|
||||
depth=12,
|
||||
num_heads=12,
|
||||
mlp_ratio=4.,
|
||||
qkv_bias=False,
|
||||
qk_scale=None,
|
||||
drop_rate=0.,
|
||||
attn_drop_rate=0.,
|
||||
drop_path_rate=0.,
|
||||
norm_layer=nn.LayerNorm,
|
||||
init_values=None,
|
||||
tubelet_size=2,
|
||||
use_learnable_pos_emb=False,
|
||||
with_cp=False,
|
||||
all_frames=16,
|
||||
cos_attn=False):
|
||||
super().__init__()
|
||||
self.num_classes = num_classes
|
||||
# num_features for consistency with other models
|
||||
self.num_features = self.embed_dim = embed_dim
|
||||
self.patch_embed = PatchEmbed(
|
||||
img_size=img_size,
|
||||
patch_size=patch_size,
|
||||
in_chans=in_chans,
|
||||
embed_dim=embed_dim,
|
||||
num_frames=all_frames,
|
||||
tubelet_size=tubelet_size)
|
||||
num_patches = self.patch_embed.num_patches
|
||||
self.with_cp = with_cp
|
||||
|
||||
if use_learnable_pos_emb:
|
||||
self.pos_embed = nn.Parameter(
|
||||
torch.zeros(1, num_patches + 1, embed_dim))
|
||||
else:
|
||||
# sine-cosine positional embeddings
|
||||
self.pos_embed = get_sinusoid_encoding_table(
|
||||
num_patches, embed_dim)
|
||||
|
||||
dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)
|
||||
] # stochastic depth decay rule
|
||||
self.blocks = nn.ModuleList([
|
||||
Block(
|
||||
dim=embed_dim,
|
||||
num_heads=num_heads,
|
||||
mlp_ratio=mlp_ratio,
|
||||
qkv_bias=qkv_bias,
|
||||
qk_scale=qk_scale,
|
||||
drop=drop_rate,
|
||||
attn_drop=attn_drop_rate,
|
||||
drop_path=dpr[i],
|
||||
norm_layer=norm_layer,
|
||||
init_values=init_values,
|
||||
cos_attn=cos_attn) for i in range(depth)
|
||||
])
|
||||
self.norm = norm_layer(embed_dim)
|
||||
self.head = nn.Linear(
|
||||
embed_dim, num_classes) if num_classes > 0 else nn.Identity()
|
||||
|
||||
if use_learnable_pos_emb:
|
||||
trunc_normal_(self.pos_embed, std=.02)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
nn.init.xavier_uniform_(m.weight)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
|
||||
def get_num_layers(self):
|
||||
return len(self.blocks)
|
||||
|
||||
@torch.jit.ignore
|
||||
def no_weight_decay(self):
|
||||
return {'pos_embed', 'cls_token'}
|
||||
|
||||
def get_classifier(self):
|
||||
return self.head
|
||||
|
||||
def reset_classifier(self, num_classes, global_pool=''):
|
||||
self.num_classes = num_classes
|
||||
self.head = nn.Linear(
|
||||
self.embed_dim, num_classes) if num_classes > 0 else nn.Identity()
|
||||
|
||||
def forward_features(self, x, mask):
|
||||
x = self.patch_embed(x)
|
||||
|
||||
x = x + self.pos_embed.type_as(x).to(x.device).clone().detach()
|
||||
|
||||
B, _, C = x.shape
|
||||
x_vis = x[~mask].reshape(B, -1, C) # ~mask means visible
|
||||
|
||||
for blk in self.blocks:
|
||||
if self.with_cp:
|
||||
x_vis = cp.checkpoint(blk, x_vis)
|
||||
else:
|
||||
x_vis = blk(x_vis)
|
||||
|
||||
x_vis = self.norm(x_vis)
|
||||
return x_vis
|
||||
|
||||
def forward(self, x, mask):
|
||||
x = self.forward_features(x, mask)
|
||||
x = self.head(x)
|
||||
return x
|
||||
|
||||
|
||||
class PretrainVisionTransformerDecoder(nn.Module):
|
||||
""" Vision Transformer with support for patch or hybrid CNN input stage
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
patch_size=16,
|
||||
num_classes=768,
|
||||
embed_dim=768,
|
||||
depth=12,
|
||||
num_heads=12,
|
||||
mlp_ratio=4.,
|
||||
qkv_bias=False,
|
||||
qk_scale=None,
|
||||
drop_rate=0.,
|
||||
attn_drop_rate=0.,
|
||||
drop_path_rate=0.,
|
||||
norm_layer=nn.LayerNorm,
|
||||
init_values=None,
|
||||
num_patches=196,
|
||||
tubelet_size=2,
|
||||
with_cp=False,
|
||||
cos_attn=False):
|
||||
super().__init__()
|
||||
self.num_classes = num_classes
|
||||
assert num_classes == 3 * tubelet_size * patch_size**2
|
||||
# num_features for consistency with other models
|
||||
self.num_features = self.embed_dim = embed_dim
|
||||
self.patch_size = patch_size
|
||||
self.with_cp = with_cp
|
||||
|
||||
dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)
|
||||
] # stochastic depth decay rule
|
||||
self.blocks = nn.ModuleList([
|
||||
Block(
|
||||
dim=embed_dim,
|
||||
num_heads=num_heads,
|
||||
mlp_ratio=mlp_ratio,
|
||||
qkv_bias=qkv_bias,
|
||||
qk_scale=qk_scale,
|
||||
drop=drop_rate,
|
||||
attn_drop=attn_drop_rate,
|
||||
drop_path=dpr[i],
|
||||
norm_layer=norm_layer,
|
||||
init_values=init_values,
|
||||
cos_attn=cos_attn) for i in range(depth)
|
||||
])
|
||||
self.norm = norm_layer(embed_dim)
|
||||
self.head = nn.Linear(
|
||||
embed_dim, num_classes) if num_classes > 0 else nn.Identity()
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
nn.init.xavier_uniform_(m.weight)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
|
||||
def get_num_layers(self):
|
||||
return len(self.blocks)
|
||||
|
||||
@torch.jit.ignore
|
||||
def no_weight_decay(self):
|
||||
return {'pos_embed', 'cls_token'}
|
||||
|
||||
def get_classifier(self):
|
||||
return self.head
|
||||
|
||||
def reset_classifier(self, num_classes, global_pool=''):
|
||||
self.num_classes = num_classes
|
||||
self.head = nn.Linear(
|
||||
self.embed_dim, num_classes) if num_classes > 0 else nn.Identity()
|
||||
|
||||
def forward(self, x, return_token_num):
|
||||
for blk in self.blocks:
|
||||
if self.with_cp:
|
||||
x = cp.checkpoint(blk, x)
|
||||
else:
|
||||
x = blk(x)
|
||||
|
||||
if return_token_num > 0:
|
||||
# only return the mask tokens predict pixels
|
||||
x = self.head(self.norm(x[:, -return_token_num:]))
|
||||
else:
|
||||
# [B, N, 3*16^2]
|
||||
x = self.head(self.norm(x))
|
||||
return x
|
||||
|
||||
|
||||
class PretrainVisionTransformer(nn.Module):
|
||||
""" Vision Transformer with support for patch or hybrid CNN input stage
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
img_size=224,
|
||||
patch_size=16,
|
||||
encoder_in_chans=3,
|
||||
encoder_num_classes=0,
|
||||
encoder_embed_dim=768,
|
||||
encoder_depth=12,
|
||||
encoder_num_heads=12,
|
||||
decoder_num_classes=1536, # decoder_num_classes=768
|
||||
decoder_embed_dim=512,
|
||||
decoder_depth=8,
|
||||
decoder_num_heads=8,
|
||||
mlp_ratio=4.,
|
||||
qkv_bias=False,
|
||||
qk_scale=None,
|
||||
drop_rate=0.,
|
||||
attn_drop_rate=0.,
|
||||
drop_path_rate=0.,
|
||||
norm_layer=nn.LayerNorm,
|
||||
init_values=0.,
|
||||
use_learnable_pos_emb=False,
|
||||
tubelet_size=2,
|
||||
num_classes=0, # avoid the error from create_fn in timm
|
||||
in_chans=0, # avoid the error from create_fn in timm
|
||||
with_cp=False,
|
||||
all_frames=16,
|
||||
cos_attn=False,
|
||||
):
|
||||
super().__init__()
|
||||
self.encoder = PretrainVisionTransformerEncoder(
|
||||
img_size=img_size,
|
||||
patch_size=patch_size,
|
||||
in_chans=encoder_in_chans,
|
||||
num_classes=encoder_num_classes,
|
||||
embed_dim=encoder_embed_dim,
|
||||
depth=encoder_depth,
|
||||
num_heads=encoder_num_heads,
|
||||
mlp_ratio=mlp_ratio,
|
||||
qkv_bias=qkv_bias,
|
||||
qk_scale=qk_scale,
|
||||
drop_rate=drop_rate,
|
||||
attn_drop_rate=attn_drop_rate,
|
||||
drop_path_rate=drop_path_rate,
|
||||
norm_layer=norm_layer,
|
||||
init_values=init_values,
|
||||
tubelet_size=tubelet_size,
|
||||
use_learnable_pos_emb=use_learnable_pos_emb,
|
||||
with_cp=with_cp,
|
||||
all_frames=all_frames,
|
||||
cos_attn=cos_attn)
|
||||
|
||||
self.decoder = PretrainVisionTransformerDecoder(
|
||||
patch_size=patch_size,
|
||||
num_patches=self.encoder.patch_embed.num_patches,
|
||||
num_classes=decoder_num_classes,
|
||||
embed_dim=decoder_embed_dim,
|
||||
depth=decoder_depth,
|
||||
num_heads=decoder_num_heads,
|
||||
mlp_ratio=mlp_ratio,
|
||||
qkv_bias=qkv_bias,
|
||||
qk_scale=qk_scale,
|
||||
drop_rate=drop_rate,
|
||||
attn_drop_rate=attn_drop_rate,
|
||||
drop_path_rate=drop_path_rate,
|
||||
norm_layer=norm_layer,
|
||||
init_values=init_values,
|
||||
tubelet_size=tubelet_size,
|
||||
with_cp=with_cp,
|
||||
cos_attn=cos_attn)
|
||||
|
||||
self.encoder_to_decoder = nn.Linear(
|
||||
encoder_embed_dim, decoder_embed_dim, bias=False)
|
||||
|
||||
self.mask_token = nn.Parameter(torch.zeros(1, 1, decoder_embed_dim))
|
||||
|
||||
self.pos_embed = get_sinusoid_encoding_table(
|
||||
self.encoder.patch_embed.num_patches, decoder_embed_dim)
|
||||
|
||||
trunc_normal_(self.mask_token, std=.02)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
nn.init.xavier_uniform_(m.weight)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
|
||||
def get_num_layers(self):
|
||||
return len(self.blocks)
|
||||
|
||||
@torch.jit.ignore
|
||||
def no_weight_decay(self):
|
||||
return {'pos_embed', 'cls_token', 'mask_token'}
|
||||
|
||||
def forward(self, x, mask, decode_mask=None):
|
||||
decode_vis = mask if decode_mask is None else ~decode_mask
|
||||
|
||||
x_vis = self.encoder(x, mask) # [B, N_vis, C_e]
|
||||
x_vis = self.encoder_to_decoder(x_vis) # [B, N_vis, C_d]
|
||||
B, N_vis, C = x_vis.shape
|
||||
|
||||
# we don't unshuffle the correct visible token order,
|
||||
# but shuffle the pos embedding accorddingly.
|
||||
expand_pos_embed = self.pos_embed.expand(B, -1, -1).type_as(x).to(
|
||||
x.device).clone().detach()
|
||||
pos_emd_vis = expand_pos_embed[~mask].reshape(B, -1, C)
|
||||
pos_emd_mask = expand_pos_embed[decode_vis].reshape(B, -1, C)
|
||||
|
||||
# [B, N, C_d]
|
||||
x_full = torch.cat(
|
||||
[x_vis + pos_emd_vis, self.mask_token + pos_emd_mask], dim=1)
|
||||
# NOTE: if N_mask==0, the shape of x is [B, N_mask, 3 * 16 * 16]
|
||||
x = self.decoder(x_full, pos_emd_mask.shape[1])
|
||||
|
||||
return x
|
||||
|
||||
|
||||
def pretrain_videomae_small_patch16_224(pretrained=False, **kwargs):
|
||||
model = PretrainVisionTransformer(
|
||||
img_size=224,
|
||||
patch_size=16,
|
||||
encoder_embed_dim=384,
|
||||
encoder_depth=12,
|
||||
encoder_num_heads=6,
|
||||
encoder_num_classes=0,
|
||||
decoder_num_classes=1536, # 16 * 16 * 3 * 2
|
||||
decoder_embed_dim=192,
|
||||
decoder_num_heads=3,
|
||||
mlp_ratio=4,
|
||||
qkv_bias=True,
|
||||
norm_layer=partial(nn.LayerNorm, eps=1e-6),
|
||||
**kwargs)
|
||||
model.default_cfg = _cfg()
|
||||
if pretrained:
|
||||
checkpoint = torch.load(kwargs["init_ckpt"], map_location="cpu")
|
||||
model.load_state_dict(checkpoint["model"])
|
||||
return model
|
||||
|
||||
|
||||
def pretrain_videomae_base_patch16_224(pretrained=False, **kwargs):
|
||||
model = PretrainVisionTransformer(
|
||||
img_size=224,
|
||||
patch_size=16,
|
||||
encoder_embed_dim=768,
|
||||
encoder_depth=12,
|
||||
encoder_num_heads=12,
|
||||
encoder_num_classes=0,
|
||||
decoder_num_classes=1536, # 16 * 16 * 3 * 2
|
||||
decoder_embed_dim=384,
|
||||
decoder_num_heads=6,
|
||||
mlp_ratio=4,
|
||||
qkv_bias=True,
|
||||
norm_layer=partial(nn.LayerNorm, eps=1e-6),
|
||||
**kwargs)
|
||||
model.default_cfg = _cfg()
|
||||
if pretrained:
|
||||
checkpoint = torch.load(kwargs["init_ckpt"], map_location="cpu")
|
||||
model.load_state_dict(checkpoint["model"])
|
||||
return model
|
||||
|
||||
|
||||
def pretrain_videomae_large_patch16_224(pretrained=False, **kwargs):
|
||||
model = PretrainVisionTransformer(
|
||||
img_size=224,
|
||||
patch_size=16,
|
||||
encoder_embed_dim=1024,
|
||||
encoder_depth=24,
|
||||
encoder_num_heads=16,
|
||||
encoder_num_classes=0,
|
||||
decoder_num_classes=1536, # 16 * 16 * 3 * 2
|
||||
decoder_embed_dim=512,
|
||||
decoder_num_heads=8,
|
||||
mlp_ratio=4,
|
||||
qkv_bias=True,
|
||||
norm_layer=partial(nn.LayerNorm, eps=1e-6),
|
||||
**kwargs)
|
||||
model.default_cfg = _cfg()
|
||||
if pretrained:
|
||||
checkpoint = torch.load(kwargs["init_ckpt"], map_location="cpu")
|
||||
model.load_state_dict(checkpoint["model"])
|
||||
return model
|
||||
|
||||
|
||||
def pretrain_videomae_huge_patch16_224(pretrained=False, **kwargs):
|
||||
model = PretrainVisionTransformer(
|
||||
img_size=224,
|
||||
patch_size=16,
|
||||
encoder_embed_dim=1280,
|
||||
encoder_depth=32,
|
||||
encoder_num_heads=16,
|
||||
encoder_num_classes=0,
|
||||
decoder_num_classes=1536, # 16 * 16 * 3 * 2
|
||||
decoder_embed_dim=512,
|
||||
decoder_num_heads=8,
|
||||
mlp_ratio=4,
|
||||
qkv_bias=True,
|
||||
norm_layer=partial(nn.LayerNorm, eps=1e-6),
|
||||
**kwargs)
|
||||
model.default_cfg = _cfg()
|
||||
if pretrained:
|
||||
checkpoint = torch.load(kwargs["init_ckpt"], map_location="cpu")
|
||||
model.load_state_dict(checkpoint["model"])
|
||||
return model
|
||||
|
||||
|
||||
def pretrain_videomae_giant_patch14_224(pretrained=False, **kwargs):
|
||||
model = PretrainVisionTransformer(
|
||||
img_size=224,
|
||||
patch_size=14,
|
||||
encoder_embed_dim=1408,
|
||||
encoder_depth=40,
|
||||
encoder_num_heads=16,
|
||||
encoder_num_classes=0,
|
||||
decoder_num_classes=1176, # 14 * 14 * 3 * 2,
|
||||
decoder_embed_dim=512,
|
||||
decoder_num_heads=8,
|
||||
mlp_ratio=48 / 11,
|
||||
qkv_bias=True,
|
||||
norm_layer=partial(nn.LayerNorm, eps=1e-6),
|
||||
**kwargs)
|
||||
model.default_cfg = _cfg()
|
||||
if pretrained:
|
||||
checkpoint = torch.load(kwargs["init_ckpt"], map_location="cpu")
|
||||
model.load_state_dict(checkpoint["model"])
|
||||
return model
|
||||
# --------------------------------------------------------
|
||||
# Based on BEiT, timm, DINO and DeiT code bases
|
||||
# https://github.com/microsoft/unilm/tree/master/beit
|
||||
# https://github.com/rwightman/pytorch-image-models/tree/master/timm
|
||||
# https://github.com/facebookresearch/deit
|
||||
# https://github.com/facebookresearch/dino
|
||||
# --------------------------------------------------------'
|
||||
from functools import partial
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.utils.checkpoint as cp
|
||||
|
||||
from .videomaev2_finetune import (
|
||||
Block,
|
||||
PatchEmbed,
|
||||
_cfg,
|
||||
get_sinusoid_encoding_table,
|
||||
)
|
||||
|
||||
from .videomaev2_finetune import trunc_normal_ as __call_trunc_normal_
|
||||
|
||||
def trunc_normal_(tensor, mean=0., std=1.):
|
||||
__call_trunc_normal_(tensor, mean=mean, std=std, a=-std, b=std)
|
||||
|
||||
|
||||
class PretrainVisionTransformerEncoder(nn.Module):
|
||||
""" Vision Transformer with support for patch or hybrid CNN input stage
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
img_size=224,
|
||||
patch_size=16,
|
||||
in_chans=3,
|
||||
num_classes=0,
|
||||
embed_dim=768,
|
||||
depth=12,
|
||||
num_heads=12,
|
||||
mlp_ratio=4.,
|
||||
qkv_bias=False,
|
||||
qk_scale=None,
|
||||
drop_rate=0.,
|
||||
attn_drop_rate=0.,
|
||||
drop_path_rate=0.,
|
||||
norm_layer=nn.LayerNorm,
|
||||
init_values=None,
|
||||
tubelet_size=2,
|
||||
use_learnable_pos_emb=False,
|
||||
with_cp=False,
|
||||
all_frames=16,
|
||||
cos_attn=False):
|
||||
super().__init__()
|
||||
self.num_classes = num_classes
|
||||
# num_features for consistency with other models
|
||||
self.num_features = self.embed_dim = embed_dim
|
||||
self.patch_embed = PatchEmbed(
|
||||
img_size=img_size,
|
||||
patch_size=patch_size,
|
||||
in_chans=in_chans,
|
||||
embed_dim=embed_dim,
|
||||
num_frames=all_frames,
|
||||
tubelet_size=tubelet_size)
|
||||
num_patches = self.patch_embed.num_patches
|
||||
self.with_cp = with_cp
|
||||
|
||||
if use_learnable_pos_emb:
|
||||
self.pos_embed = nn.Parameter(
|
||||
torch.zeros(1, num_patches + 1, embed_dim))
|
||||
else:
|
||||
# sine-cosine positional embeddings
|
||||
self.pos_embed = get_sinusoid_encoding_table(
|
||||
num_patches, embed_dim)
|
||||
|
||||
dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)
|
||||
] # stochastic depth decay rule
|
||||
self.blocks = nn.ModuleList([
|
||||
Block(
|
||||
dim=embed_dim,
|
||||
num_heads=num_heads,
|
||||
mlp_ratio=mlp_ratio,
|
||||
qkv_bias=qkv_bias,
|
||||
qk_scale=qk_scale,
|
||||
drop=drop_rate,
|
||||
attn_drop=attn_drop_rate,
|
||||
drop_path=dpr[i],
|
||||
norm_layer=norm_layer,
|
||||
init_values=init_values,
|
||||
cos_attn=cos_attn) for i in range(depth)
|
||||
])
|
||||
self.norm = norm_layer(embed_dim)
|
||||
self.head = nn.Linear(
|
||||
embed_dim, num_classes) if num_classes > 0 else nn.Identity()
|
||||
|
||||
if use_learnable_pos_emb:
|
||||
trunc_normal_(self.pos_embed, std=.02)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
nn.init.xavier_uniform_(m.weight)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
|
||||
def get_num_layers(self):
|
||||
return len(self.blocks)
|
||||
|
||||
@torch.jit.ignore
|
||||
def no_weight_decay(self):
|
||||
return {'pos_embed', 'cls_token'}
|
||||
|
||||
def get_classifier(self):
|
||||
return self.head
|
||||
|
||||
def reset_classifier(self, num_classes, global_pool=''):
|
||||
self.num_classes = num_classes
|
||||
self.head = nn.Linear(
|
||||
self.embed_dim, num_classes) if num_classes > 0 else nn.Identity()
|
||||
|
||||
def forward_features(self, x, mask):
|
||||
x = self.patch_embed(x)
|
||||
|
||||
x = x + self.pos_embed.type_as(x).to(x.device).clone().detach()
|
||||
|
||||
B, _, C = x.shape
|
||||
x_vis = x[~mask].reshape(B, -1, C) # ~mask means visible
|
||||
|
||||
for blk in self.blocks:
|
||||
if self.with_cp:
|
||||
x_vis = cp.checkpoint(blk, x_vis)
|
||||
else:
|
||||
x_vis = blk(x_vis)
|
||||
|
||||
x_vis = self.norm(x_vis)
|
||||
return x_vis
|
||||
|
||||
def forward(self, x, mask):
|
||||
x = self.forward_features(x, mask)
|
||||
x = self.head(x)
|
||||
return x
|
||||
|
||||
|
||||
class PretrainVisionTransformerDecoder(nn.Module):
|
||||
""" Vision Transformer with support for patch or hybrid CNN input stage
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
patch_size=16,
|
||||
num_classes=768,
|
||||
embed_dim=768,
|
||||
depth=12,
|
||||
num_heads=12,
|
||||
mlp_ratio=4.,
|
||||
qkv_bias=False,
|
||||
qk_scale=None,
|
||||
drop_rate=0.,
|
||||
attn_drop_rate=0.,
|
||||
drop_path_rate=0.,
|
||||
norm_layer=nn.LayerNorm,
|
||||
init_values=None,
|
||||
num_patches=196,
|
||||
tubelet_size=2,
|
||||
with_cp=False,
|
||||
cos_attn=False):
|
||||
super().__init__()
|
||||
self.num_classes = num_classes
|
||||
assert num_classes == 3 * tubelet_size * patch_size**2
|
||||
# num_features for consistency with other models
|
||||
self.num_features = self.embed_dim = embed_dim
|
||||
self.patch_size = patch_size
|
||||
self.with_cp = with_cp
|
||||
|
||||
dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)
|
||||
] # stochastic depth decay rule
|
||||
self.blocks = nn.ModuleList([
|
||||
Block(
|
||||
dim=embed_dim,
|
||||
num_heads=num_heads,
|
||||
mlp_ratio=mlp_ratio,
|
||||
qkv_bias=qkv_bias,
|
||||
qk_scale=qk_scale,
|
||||
drop=drop_rate,
|
||||
attn_drop=attn_drop_rate,
|
||||
drop_path=dpr[i],
|
||||
norm_layer=norm_layer,
|
||||
init_values=init_values,
|
||||
cos_attn=cos_attn) for i in range(depth)
|
||||
])
|
||||
self.norm = norm_layer(embed_dim)
|
||||
self.head = nn.Linear(
|
||||
embed_dim, num_classes) if num_classes > 0 else nn.Identity()
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
nn.init.xavier_uniform_(m.weight)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
|
||||
def get_num_layers(self):
|
||||
return len(self.blocks)
|
||||
|
||||
@torch.jit.ignore
|
||||
def no_weight_decay(self):
|
||||
return {'pos_embed', 'cls_token'}
|
||||
|
||||
def get_classifier(self):
|
||||
return self.head
|
||||
|
||||
def reset_classifier(self, num_classes, global_pool=''):
|
||||
self.num_classes = num_classes
|
||||
self.head = nn.Linear(
|
||||
self.embed_dim, num_classes) if num_classes > 0 else nn.Identity()
|
||||
|
||||
def forward(self, x, return_token_num):
|
||||
for blk in self.blocks:
|
||||
if self.with_cp:
|
||||
x = cp.checkpoint(blk, x)
|
||||
else:
|
||||
x = blk(x)
|
||||
|
||||
if return_token_num > 0:
|
||||
# only return the mask tokens predict pixels
|
||||
x = self.head(self.norm(x[:, -return_token_num:]))
|
||||
else:
|
||||
# [B, N, 3*16^2]
|
||||
x = self.head(self.norm(x))
|
||||
return x
|
||||
|
||||
|
||||
class PretrainVisionTransformer(nn.Module):
|
||||
""" Vision Transformer with support for patch or hybrid CNN input stage
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
img_size=224,
|
||||
patch_size=16,
|
||||
encoder_in_chans=3,
|
||||
encoder_num_classes=0,
|
||||
encoder_embed_dim=768,
|
||||
encoder_depth=12,
|
||||
encoder_num_heads=12,
|
||||
decoder_num_classes=1536, # decoder_num_classes=768
|
||||
decoder_embed_dim=512,
|
||||
decoder_depth=8,
|
||||
decoder_num_heads=8,
|
||||
mlp_ratio=4.,
|
||||
qkv_bias=False,
|
||||
qk_scale=None,
|
||||
drop_rate=0.,
|
||||
attn_drop_rate=0.,
|
||||
drop_path_rate=0.,
|
||||
norm_layer=nn.LayerNorm,
|
||||
init_values=0.,
|
||||
use_learnable_pos_emb=False,
|
||||
tubelet_size=2,
|
||||
num_classes=0, # avoid the error from create_fn in timm
|
||||
in_chans=0, # avoid the error from create_fn in timm
|
||||
with_cp=False,
|
||||
all_frames=16,
|
||||
cos_attn=False,
|
||||
):
|
||||
super().__init__()
|
||||
self.encoder = PretrainVisionTransformerEncoder(
|
||||
img_size=img_size,
|
||||
patch_size=patch_size,
|
||||
in_chans=encoder_in_chans,
|
||||
num_classes=encoder_num_classes,
|
||||
embed_dim=encoder_embed_dim,
|
||||
depth=encoder_depth,
|
||||
num_heads=encoder_num_heads,
|
||||
mlp_ratio=mlp_ratio,
|
||||
qkv_bias=qkv_bias,
|
||||
qk_scale=qk_scale,
|
||||
drop_rate=drop_rate,
|
||||
attn_drop_rate=attn_drop_rate,
|
||||
drop_path_rate=drop_path_rate,
|
||||
norm_layer=norm_layer,
|
||||
init_values=init_values,
|
||||
tubelet_size=tubelet_size,
|
||||
use_learnable_pos_emb=use_learnable_pos_emb,
|
||||
with_cp=with_cp,
|
||||
all_frames=all_frames,
|
||||
cos_attn=cos_attn)
|
||||
|
||||
self.decoder = PretrainVisionTransformerDecoder(
|
||||
patch_size=patch_size,
|
||||
num_patches=self.encoder.patch_embed.num_patches,
|
||||
num_classes=decoder_num_classes,
|
||||
embed_dim=decoder_embed_dim,
|
||||
depth=decoder_depth,
|
||||
num_heads=decoder_num_heads,
|
||||
mlp_ratio=mlp_ratio,
|
||||
qkv_bias=qkv_bias,
|
||||
qk_scale=qk_scale,
|
||||
drop_rate=drop_rate,
|
||||
attn_drop_rate=attn_drop_rate,
|
||||
drop_path_rate=drop_path_rate,
|
||||
norm_layer=norm_layer,
|
||||
init_values=init_values,
|
||||
tubelet_size=tubelet_size,
|
||||
with_cp=with_cp,
|
||||
cos_attn=cos_attn)
|
||||
|
||||
self.encoder_to_decoder = nn.Linear(
|
||||
encoder_embed_dim, decoder_embed_dim, bias=False)
|
||||
|
||||
self.mask_token = nn.Parameter(torch.zeros(1, 1, decoder_embed_dim))
|
||||
|
||||
self.pos_embed = get_sinusoid_encoding_table(
|
||||
self.encoder.patch_embed.num_patches, decoder_embed_dim)
|
||||
|
||||
trunc_normal_(self.mask_token, std=.02)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
nn.init.xavier_uniform_(m.weight)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
|
||||
def get_num_layers(self):
|
||||
return len(self.blocks)
|
||||
|
||||
@torch.jit.ignore
|
||||
def no_weight_decay(self):
|
||||
return {'pos_embed', 'cls_token', 'mask_token'}
|
||||
|
||||
def forward(self, x, mask, decode_mask=None):
|
||||
decode_vis = mask if decode_mask is None else ~decode_mask
|
||||
|
||||
x_vis = self.encoder(x, mask) # [B, N_vis, C_e]
|
||||
x_vis = self.encoder_to_decoder(x_vis) # [B, N_vis, C_d]
|
||||
B, N_vis, C = x_vis.shape
|
||||
|
||||
# we don't unshuffle the correct visible token order,
|
||||
# but shuffle the pos embedding accorddingly.
|
||||
expand_pos_embed = self.pos_embed.expand(B, -1, -1).type_as(x).to(
|
||||
x.device).clone().detach()
|
||||
pos_emd_vis = expand_pos_embed[~mask].reshape(B, -1, C)
|
||||
pos_emd_mask = expand_pos_embed[decode_vis].reshape(B, -1, C)
|
||||
|
||||
# [B, N, C_d]
|
||||
x_full = torch.cat(
|
||||
[x_vis + pos_emd_vis, self.mask_token + pos_emd_mask], dim=1)
|
||||
# NOTE: if N_mask==0, the shape of x is [B, N_mask, 3 * 16 * 16]
|
||||
x = self.decoder(x_full, pos_emd_mask.shape[1])
|
||||
|
||||
return x
|
||||
|
||||
|
||||
def pretrain_videomae_small_patch16_224(pretrained=False, **kwargs):
|
||||
model = PretrainVisionTransformer(
|
||||
img_size=224,
|
||||
patch_size=16,
|
||||
encoder_embed_dim=384,
|
||||
encoder_depth=12,
|
||||
encoder_num_heads=6,
|
||||
encoder_num_classes=0,
|
||||
decoder_num_classes=1536, # 16 * 16 * 3 * 2
|
||||
decoder_embed_dim=192,
|
||||
decoder_num_heads=3,
|
||||
mlp_ratio=4,
|
||||
qkv_bias=True,
|
||||
norm_layer=partial(nn.LayerNorm, eps=1e-6),
|
||||
**kwargs)
|
||||
model.default_cfg = _cfg()
|
||||
if pretrained:
|
||||
checkpoint = torch.load(kwargs["init_ckpt"], map_location="cpu")
|
||||
model.load_state_dict(checkpoint["model"])
|
||||
return model
|
||||
|
||||
|
||||
def pretrain_videomae_base_patch16_224(pretrained=False, **kwargs):
|
||||
model = PretrainVisionTransformer(
|
||||
img_size=224,
|
||||
patch_size=16,
|
||||
encoder_embed_dim=768,
|
||||
encoder_depth=12,
|
||||
encoder_num_heads=12,
|
||||
encoder_num_classes=0,
|
||||
decoder_num_classes=1536, # 16 * 16 * 3 * 2
|
||||
decoder_embed_dim=384,
|
||||
decoder_num_heads=6,
|
||||
mlp_ratio=4,
|
||||
qkv_bias=True,
|
||||
norm_layer=partial(nn.LayerNorm, eps=1e-6),
|
||||
**kwargs)
|
||||
model.default_cfg = _cfg()
|
||||
if pretrained:
|
||||
checkpoint = torch.load(kwargs["init_ckpt"], map_location="cpu")
|
||||
model.load_state_dict(checkpoint["model"])
|
||||
return model
|
||||
|
||||
|
||||
def pretrain_videomae_large_patch16_224(pretrained=False, **kwargs):
|
||||
model = PretrainVisionTransformer(
|
||||
img_size=224,
|
||||
patch_size=16,
|
||||
encoder_embed_dim=1024,
|
||||
encoder_depth=24,
|
||||
encoder_num_heads=16,
|
||||
encoder_num_classes=0,
|
||||
decoder_num_classes=1536, # 16 * 16 * 3 * 2
|
||||
decoder_embed_dim=512,
|
||||
decoder_num_heads=8,
|
||||
mlp_ratio=4,
|
||||
qkv_bias=True,
|
||||
norm_layer=partial(nn.LayerNorm, eps=1e-6),
|
||||
**kwargs)
|
||||
model.default_cfg = _cfg()
|
||||
if pretrained:
|
||||
checkpoint = torch.load(kwargs["init_ckpt"], map_location="cpu")
|
||||
model.load_state_dict(checkpoint["model"])
|
||||
return model
|
||||
|
||||
|
||||
def pretrain_videomae_huge_patch16_224(pretrained=False, **kwargs):
|
||||
model = PretrainVisionTransformer(
|
||||
img_size=224,
|
||||
patch_size=16,
|
||||
encoder_embed_dim=1280,
|
||||
encoder_depth=32,
|
||||
encoder_num_heads=16,
|
||||
encoder_num_classes=0,
|
||||
decoder_num_classes=1536, # 16 * 16 * 3 * 2
|
||||
decoder_embed_dim=512,
|
||||
decoder_num_heads=8,
|
||||
mlp_ratio=4,
|
||||
qkv_bias=True,
|
||||
norm_layer=partial(nn.LayerNorm, eps=1e-6),
|
||||
**kwargs)
|
||||
model.default_cfg = _cfg()
|
||||
if pretrained:
|
||||
checkpoint = torch.load(kwargs["init_ckpt"], map_location="cpu")
|
||||
model.load_state_dict(checkpoint["model"])
|
||||
return model
|
||||
|
||||
|
||||
def pretrain_videomae_giant_patch14_224(pretrained=False, **kwargs):
|
||||
model = PretrainVisionTransformer(
|
||||
img_size=224,
|
||||
patch_size=14,
|
||||
encoder_embed_dim=1408,
|
||||
encoder_depth=40,
|
||||
encoder_num_heads=16,
|
||||
encoder_num_classes=0,
|
||||
decoder_num_classes=1176, # 14 * 14 * 3 * 2,
|
||||
decoder_embed_dim=512,
|
||||
decoder_num_heads=8,
|
||||
mlp_ratio=48 / 11,
|
||||
qkv_bias=True,
|
||||
norm_layer=partial(nn.LayerNorm, eps=1e-6),
|
||||
**kwargs)
|
||||
model.default_cfg = _cfg()
|
||||
if pretrained:
|
||||
checkpoint = torch.load(kwargs["init_ckpt"], map_location="cpu")
|
||||
model.load_state_dict(checkpoint["model"])
|
||||
return model
|
||||
|
||||
@@ -1,321 +1,321 @@
|
||||
import os
|
||||
import math
|
||||
import os.path as osp
|
||||
import random
|
||||
import pickle
|
||||
import warnings
|
||||
|
||||
import glob
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
import torch
|
||||
import torch.utils.data as data
|
||||
import torch.nn.functional as F
|
||||
import torch.distributed as dist
|
||||
from torchvision.datasets.video_utils import VideoClips
|
||||
|
||||
IMG_EXTENSIONS = ['.jpg', '.JPG', '.jpeg', '.JPEG', '.png', '.PNG']
|
||||
VID_EXTENSIONS = ['.avi', '.mp4', '.webm', '.mov', '.mkv', '.m4v']
|
||||
|
||||
|
||||
def get_dataloader(data_path, image_folder, resolution=128, sequence_length=16, sample_every_n_frames=1,
|
||||
batch_size=16, num_workers=8):
|
||||
data = VideoData(data_path, image_folder, resolution, sequence_length, sample_every_n_frames, batch_size, num_workers)
|
||||
loader = data._dataloader()
|
||||
return loader
|
||||
|
||||
|
||||
def is_image_file(filename):
|
||||
return any(filename.endswith(extension) for extension in IMG_EXTENSIONS)
|
||||
|
||||
|
||||
def get_parent_dir(path):
|
||||
return osp.basename(osp.dirname(path))
|
||||
|
||||
|
||||
def preprocess(video, resolution, sequence_length=None, in_channels=3, sample_every_n_frames=1):
|
||||
# video: THWC, {0, ..., 255}
|
||||
assert in_channels == 3
|
||||
video = video.permute(0, 3, 1, 2).float() / 255. # TCHW
|
||||
t, c, h, w = video.shape
|
||||
|
||||
# temporal crop
|
||||
if sequence_length is not None:
|
||||
assert sequence_length <= t
|
||||
video = video[:sequence_length]
|
||||
|
||||
# skip frames
|
||||
if sample_every_n_frames > 1:
|
||||
video = video[::sample_every_n_frames]
|
||||
|
||||
# scale shorter side to resolution
|
||||
scale = resolution / min(h, w)
|
||||
if h < w:
|
||||
target_size = (resolution, math.ceil(w * scale))
|
||||
else:
|
||||
target_size = (math.ceil(h * scale), resolution)
|
||||
video = F.interpolate(video, size=target_size, mode='bilinear',
|
||||
align_corners=False, antialias=True)
|
||||
|
||||
# center crop
|
||||
t, c, h, w = video.shape
|
||||
w_start = (w - resolution) // 2
|
||||
h_start = (h - resolution) // 2
|
||||
video = video[:, :, h_start:h_start + resolution, w_start:w_start + resolution]
|
||||
video = video.permute(1, 0, 2, 3).contiguous() # CTHW
|
||||
|
||||
return {'video': video}
|
||||
|
||||
|
||||
def preprocess_image(image):
|
||||
# [0, 1] => [-1, 1]
|
||||
img = torch.from_numpy(image)
|
||||
return img
|
||||
|
||||
|
||||
class VideoData(data.Dataset):
|
||||
""" Class to create dataloaders for video datasets
|
||||
|
||||
Args:
|
||||
data_path: Path to the folder with video frames or videos.
|
||||
image_folder: If True, the data is stored as images in folders.
|
||||
resolution: Resolution of the returned videos.
|
||||
sequence_length: Length of extracted video sequences.
|
||||
sample_every_n_frames: Sample every n frames from the video.
|
||||
batch_size: Batch size.
|
||||
num_workers: Number of workers for the dataloader.
|
||||
shuffle: If True, shuffle the data.
|
||||
"""
|
||||
|
||||
def __init__(self, data_path: str, image_folder: bool, resolution: int, sequence_length: int,
|
||||
sample_every_n_frames: int, batch_size: int, num_workers: int, shuffle: bool = True):
|
||||
super().__init__()
|
||||
self.data_path = data_path
|
||||
self.image_folder = image_folder
|
||||
self.resolution = resolution
|
||||
self.sequence_length = sequence_length
|
||||
self.sample_every_n_frames = sample_every_n_frames
|
||||
self.batch_size = batch_size
|
||||
self.num_workers = num_workers
|
||||
self.shuffle = shuffle
|
||||
|
||||
def _dataset(self):
|
||||
'''
|
||||
Initializes and return the dataset.
|
||||
'''
|
||||
if self.image_folder:
|
||||
Dataset = FrameDataset
|
||||
dataset = Dataset(self.data_path, self.sequence_length,
|
||||
resolution=self.resolution, sample_every_n_frames=self.sample_every_n_frames)
|
||||
else:
|
||||
Dataset = VideoDataset
|
||||
dataset = Dataset(self.data_path, self.sequence_length,
|
||||
resolution=self.resolution, sample_every_n_frames=self.sample_every_n_frames)
|
||||
return dataset
|
||||
|
||||
def _dataloader(self):
|
||||
'''
|
||||
Initializes and returns the dataloader.
|
||||
'''
|
||||
dataset = self._dataset()
|
||||
if dist.is_initialized():
|
||||
sampler = data.distributed.DistributedSampler(
|
||||
dataset, num_replicas=dist.get_world_size(), rank=dist.get_rank()
|
||||
)
|
||||
else:
|
||||
sampler = None
|
||||
dataloader = data.DataLoader(
|
||||
dataset,
|
||||
batch_size=self.batch_size,
|
||||
num_workers=self.num_workers,
|
||||
pin_memory=True,
|
||||
sampler=sampler,
|
||||
shuffle=sampler is None and self.shuffle is True
|
||||
)
|
||||
return dataloader
|
||||
|
||||
|
||||
class VideoDataset(data.Dataset):
|
||||
"""
|
||||
Generic dataset for videos files stored in folders.
|
||||
Videos of the same class are expected to be stored in a single folder. Multiple folders can exist in the provided directory.
|
||||
The class depends on `torchvision.datasets.video_utils.VideoClips` to load the videos.
|
||||
Returns BCTHW videos in the range [0, 1].
|
||||
|
||||
Args:
|
||||
data_folder: Path to the folder with corresponding videos stored.
|
||||
sequence_length: Length of extracted video sequences.
|
||||
resolution: Resolution of the returned videos.
|
||||
sample_every_n_frames: Sample every n frames from the video.
|
||||
"""
|
||||
|
||||
def __init__(self, data_folder: str, sequence_length: int = 16, resolution: int = 128, sample_every_n_frames: int = 1):
|
||||
super().__init__()
|
||||
self.sequence_length = sequence_length
|
||||
self.resolution = resolution
|
||||
self.sample_every_n_frames = sample_every_n_frames
|
||||
|
||||
folder = data_folder
|
||||
files = sum([glob.glob(osp.join(folder, '**', f'*{ext}'), recursive=True)
|
||||
for ext in VID_EXTENSIONS], [])
|
||||
|
||||
warnings.filterwarnings('ignore')
|
||||
cache_file = osp.join(folder, f"metadata_{sequence_length}.pkl")
|
||||
if not osp.exists(cache_file):
|
||||
clips = VideoClips(files, sequence_length, num_workers=4)
|
||||
try:
|
||||
pickle.dump(clips.metadata, open(cache_file, 'wb'))
|
||||
except:
|
||||
print(f"Failed to save metadata to {cache_file}")
|
||||
else:
|
||||
metadata = pickle.load(open(cache_file, 'rb'))
|
||||
clips = VideoClips(files, sequence_length,
|
||||
_precomputed_metadata=metadata)
|
||||
|
||||
self._clips = clips
|
||||
# instead of uniformly sampling from all possible clips, we sample uniformly from all possible videos
|
||||
self._clips.get_clip_location = self.get_random_clip_from_video
|
||||
|
||||
def get_random_clip_from_video(self, idx: int) -> tuple:
|
||||
'''
|
||||
Sample a random clip starting index from the video.
|
||||
|
||||
Args:
|
||||
idx: Index of the video.
|
||||
'''
|
||||
# Note that some videos may not contain enough frames, we skip those videos here.
|
||||
while self._clips.clips[idx].shape[0] <= 0:
|
||||
idx += 1
|
||||
n_clip = self._clips.clips[idx].shape[0]
|
||||
clip_id = random.randint(0, n_clip - 1)
|
||||
return idx, clip_id
|
||||
|
||||
def __len__(self):
|
||||
return self._clips.num_videos()
|
||||
|
||||
def __getitem__(self, idx):
|
||||
resolution = self.resolution
|
||||
while True:
|
||||
try:
|
||||
video, _, _, idx = self._clips.get_clip(idx)
|
||||
except Exception as e:
|
||||
print(idx, e)
|
||||
idx = (idx + 1) % self._clips.num_clips()
|
||||
continue
|
||||
break
|
||||
|
||||
return dict(**preprocess(video, resolution, sample_every_n_frames=self.sample_every_n_frames))
|
||||
|
||||
|
||||
class FrameDataset(data.Dataset):
|
||||
"""
|
||||
Generic dataset for videos stored as images. The loading will iterates over all the folders and subfolders
|
||||
in the provided directory. Each leaf folder is assumed to contain frames from a single video.
|
||||
|
||||
Args:
|
||||
data_folder: path to the folder with video frames. The folder
|
||||
should contain folders with frames from each video.
|
||||
sequence_length: length of extracted video sequences
|
||||
resolution: resolution of the returned videos
|
||||
sample_every_n_frames: sample every n frames from the video
|
||||
"""
|
||||
|
||||
def __init__(self, data_folder, sequence_length, resolution=64, sample_every_n_frames=1):
|
||||
self.resolution = resolution
|
||||
self.sequence_length = sequence_length
|
||||
self.sample_every_n_frames = sample_every_n_frames
|
||||
self.data_all = self.load_video_frames(data_folder)
|
||||
self.video_num = len(self.data_all)
|
||||
|
||||
def __getitem__(self, index):
|
||||
batch_data = self.getTensor(index)
|
||||
return_list = {'video': batch_data}
|
||||
|
||||
return return_list
|
||||
|
||||
def load_video_frames(self, dataroot: str) -> list:
|
||||
'''
|
||||
Loads all the video frames under the dataroot and returns a list of all the video frames.
|
||||
|
||||
Args:
|
||||
dataroot: The root directory containing the video frames.
|
||||
|
||||
Returns:
|
||||
A list of all the video frames.
|
||||
|
||||
'''
|
||||
data_all = []
|
||||
frame_list = os.walk(dataroot)
|
||||
for _, meta in enumerate(frame_list):
|
||||
root = meta[0]
|
||||
try:
|
||||
frames = sorted(meta[2], key=lambda item: int(item.split('.')[0].split('_')[-1]))
|
||||
except:
|
||||
print(meta[0], meta[2])
|
||||
if len(frames) < max(0, self.sequence_length * self.sample_every_n_frames):
|
||||
continue
|
||||
frames = [
|
||||
os.path.join(root, item) for item in frames
|
||||
if is_image_file(item)
|
||||
]
|
||||
if len(frames) > max(0, self.sequence_length * self.sample_every_n_frames):
|
||||
data_all.append(frames)
|
||||
|
||||
return data_all
|
||||
|
||||
def getTensor(self, index: int) -> torch.Tensor:
|
||||
'''
|
||||
Returns a tensor of the video frames at the given index.
|
||||
|
||||
Args:
|
||||
index: The index of the video frames to return.
|
||||
|
||||
Returns:
|
||||
A BCTHW tensor in the range `[0, 1]` of the video frames at the given index.
|
||||
|
||||
'''
|
||||
video = self.data_all[index]
|
||||
video_len = len(video)
|
||||
|
||||
# load the entire video when sequence_length = -1, whiel the sample_every_n_frames has to be 1
|
||||
if self.sequence_length == -1:
|
||||
assert self.sample_every_n_frames == 1
|
||||
start_idx = 0
|
||||
end_idx = video_len
|
||||
else:
|
||||
n_frames_interval = self.sequence_length * self.sample_every_n_frames
|
||||
start_idx = random.randint(0, video_len - n_frames_interval)
|
||||
end_idx = start_idx + n_frames_interval
|
||||
img = Image.open(video[0])
|
||||
h, w = img.height, img.width
|
||||
|
||||
if h > w:
|
||||
half = (h - w) // 2
|
||||
cropsize = (0, half, w, half + w) # left, upper, right, lower
|
||||
elif w > h:
|
||||
half = (w - h) // 2
|
||||
cropsize = (half, 0, half + h, h)
|
||||
|
||||
images = []
|
||||
for i in range(start_idx, end_idx,
|
||||
self.sample_every_n_frames):
|
||||
path = video[i]
|
||||
img = Image.open(path)
|
||||
|
||||
if h != w:
|
||||
img = img.crop(cropsize)
|
||||
|
||||
img = img.resize(
|
||||
(self.resolution, self.resolution),
|
||||
Image.ANTIALIAS)
|
||||
img = np.asarray(img, dtype=np.float32)
|
||||
img /= 255.
|
||||
img_tensor = preprocess_image(img).unsqueeze(0)
|
||||
images.append(img_tensor)
|
||||
|
||||
video_clip = torch.cat(images).permute(3, 0, 1, 2)
|
||||
return video_clip
|
||||
|
||||
def __len__(self):
|
||||
return self.video_num
|
||||
import os
|
||||
import math
|
||||
import os.path as osp
|
||||
import random
|
||||
import pickle
|
||||
import warnings
|
||||
|
||||
import glob
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
import torch
|
||||
import torch.utils.data as data
|
||||
import torch.nn.functional as F
|
||||
import torch.distributed as dist
|
||||
from torchvision.datasets.video_utils import VideoClips
|
||||
|
||||
IMG_EXTENSIONS = ['.jpg', '.JPG', '.jpeg', '.JPEG', '.png', '.PNG']
|
||||
VID_EXTENSIONS = ['.avi', '.mp4', '.webm', '.mov', '.mkv', '.m4v']
|
||||
|
||||
|
||||
def get_dataloader(data_path, image_folder, resolution=128, sequence_length=16, sample_every_n_frames=1,
|
||||
batch_size=16, num_workers=8):
|
||||
data = VideoData(data_path, image_folder, resolution, sequence_length, sample_every_n_frames, batch_size, num_workers)
|
||||
loader = data._dataloader()
|
||||
return loader
|
||||
|
||||
|
||||
def is_image_file(filename):
|
||||
return any(filename.endswith(extension) for extension in IMG_EXTENSIONS)
|
||||
|
||||
|
||||
def get_parent_dir(path):
|
||||
return osp.basename(osp.dirname(path))
|
||||
|
||||
|
||||
def preprocess(video, resolution, sequence_length=None, in_channels=3, sample_every_n_frames=1):
|
||||
# video: THWC, {0, ..., 255}
|
||||
assert in_channels == 3
|
||||
video = video.permute(0, 3, 1, 2).float() / 255. # TCHW
|
||||
t, c, h, w = video.shape
|
||||
|
||||
# temporal crop
|
||||
if sequence_length is not None:
|
||||
assert sequence_length <= t
|
||||
video = video[:sequence_length]
|
||||
|
||||
# skip frames
|
||||
if sample_every_n_frames > 1:
|
||||
video = video[::sample_every_n_frames]
|
||||
|
||||
# scale shorter side to resolution
|
||||
scale = resolution / min(h, w)
|
||||
if h < w:
|
||||
target_size = (resolution, math.ceil(w * scale))
|
||||
else:
|
||||
target_size = (math.ceil(h * scale), resolution)
|
||||
video = F.interpolate(video, size=target_size, mode='bilinear',
|
||||
align_corners=False, antialias=True)
|
||||
|
||||
# center crop
|
||||
t, c, h, w = video.shape
|
||||
w_start = (w - resolution) // 2
|
||||
h_start = (h - resolution) // 2
|
||||
video = video[:, :, h_start:h_start + resolution, w_start:w_start + resolution]
|
||||
video = video.permute(1, 0, 2, 3).contiguous() # CTHW
|
||||
|
||||
return {'video': video}
|
||||
|
||||
|
||||
def preprocess_image(image):
|
||||
# [0, 1] => [-1, 1]
|
||||
img = torch.from_numpy(image)
|
||||
return img
|
||||
|
||||
|
||||
class VideoData(data.Dataset):
|
||||
""" Class to create dataloaders for video datasets
|
||||
|
||||
Args:
|
||||
data_path: Path to the folder with video frames or videos.
|
||||
image_folder: If True, the data is stored as images in folders.
|
||||
resolution: Resolution of the returned videos.
|
||||
sequence_length: Length of extracted video sequences.
|
||||
sample_every_n_frames: Sample every n frames from the video.
|
||||
batch_size: Batch size.
|
||||
num_workers: Number of workers for the dataloader.
|
||||
shuffle: If True, shuffle the data.
|
||||
"""
|
||||
|
||||
def __init__(self, data_path: str, image_folder: bool, resolution: int, sequence_length: int,
|
||||
sample_every_n_frames: int, batch_size: int, num_workers: int, shuffle: bool = True):
|
||||
super().__init__()
|
||||
self.data_path = data_path
|
||||
self.image_folder = image_folder
|
||||
self.resolution = resolution
|
||||
self.sequence_length = sequence_length
|
||||
self.sample_every_n_frames = sample_every_n_frames
|
||||
self.batch_size = batch_size
|
||||
self.num_workers = num_workers
|
||||
self.shuffle = shuffle
|
||||
|
||||
def _dataset(self):
|
||||
'''
|
||||
Initializes and return the dataset.
|
||||
'''
|
||||
if self.image_folder:
|
||||
Dataset = FrameDataset
|
||||
dataset = Dataset(self.data_path, self.sequence_length,
|
||||
resolution=self.resolution, sample_every_n_frames=self.sample_every_n_frames)
|
||||
else:
|
||||
Dataset = VideoDataset
|
||||
dataset = Dataset(self.data_path, self.sequence_length,
|
||||
resolution=self.resolution, sample_every_n_frames=self.sample_every_n_frames)
|
||||
return dataset
|
||||
|
||||
def _dataloader(self):
|
||||
'''
|
||||
Initializes and returns the dataloader.
|
||||
'''
|
||||
dataset = self._dataset()
|
||||
if dist.is_initialized():
|
||||
sampler = data.distributed.DistributedSampler(
|
||||
dataset, num_replicas=dist.get_world_size(), rank=dist.get_rank()
|
||||
)
|
||||
else:
|
||||
sampler = None
|
||||
dataloader = data.DataLoader(
|
||||
dataset,
|
||||
batch_size=self.batch_size,
|
||||
num_workers=self.num_workers,
|
||||
pin_memory=True,
|
||||
sampler=sampler,
|
||||
shuffle=sampler is None and self.shuffle is True
|
||||
)
|
||||
return dataloader
|
||||
|
||||
|
||||
class VideoDataset(data.Dataset):
|
||||
"""
|
||||
Generic dataset for videos files stored in folders.
|
||||
Videos of the same class are expected to be stored in a single folder. Multiple folders can exist in the provided directory.
|
||||
The class depends on `torchvision.datasets.video_utils.VideoClips` to load the videos.
|
||||
Returns BCTHW videos in the range [0, 1].
|
||||
|
||||
Args:
|
||||
data_folder: Path to the folder with corresponding videos stored.
|
||||
sequence_length: Length of extracted video sequences.
|
||||
resolution: Resolution of the returned videos.
|
||||
sample_every_n_frames: Sample every n frames from the video.
|
||||
"""
|
||||
|
||||
def __init__(self, data_folder: str, sequence_length: int = 16, resolution: int = 128, sample_every_n_frames: int = 1):
|
||||
super().__init__()
|
||||
self.sequence_length = sequence_length
|
||||
self.resolution = resolution
|
||||
self.sample_every_n_frames = sample_every_n_frames
|
||||
|
||||
folder = data_folder
|
||||
files = sum([glob.glob(osp.join(folder, '**', f'*{ext}'), recursive=True)
|
||||
for ext in VID_EXTENSIONS], [])
|
||||
|
||||
warnings.filterwarnings('ignore')
|
||||
cache_file = osp.join(folder, f"metadata_{sequence_length}.pkl")
|
||||
if not osp.exists(cache_file):
|
||||
clips = VideoClips(files, sequence_length, num_workers=4)
|
||||
try:
|
||||
pickle.dump(clips.metadata, open(cache_file, 'wb'))
|
||||
except:
|
||||
print(f"Failed to save metadata to {cache_file}")
|
||||
else:
|
||||
metadata = pickle.load(open(cache_file, 'rb'))
|
||||
clips = VideoClips(files, sequence_length,
|
||||
_precomputed_metadata=metadata)
|
||||
|
||||
self._clips = clips
|
||||
# instead of uniformly sampling from all possible clips, we sample uniformly from all possible videos
|
||||
self._clips.get_clip_location = self.get_random_clip_from_video
|
||||
|
||||
def get_random_clip_from_video(self, idx: int) -> tuple:
|
||||
'''
|
||||
Sample a random clip starting index from the video.
|
||||
|
||||
Args:
|
||||
idx: Index of the video.
|
||||
'''
|
||||
# Note that some videos may not contain enough frames, we skip those videos here.
|
||||
while self._clips.clips[idx].shape[0] <= 0:
|
||||
idx += 1
|
||||
n_clip = self._clips.clips[idx].shape[0]
|
||||
clip_id = random.randint(0, n_clip - 1)
|
||||
return idx, clip_id
|
||||
|
||||
def __len__(self):
|
||||
return self._clips.num_videos()
|
||||
|
||||
def __getitem__(self, idx):
|
||||
resolution = self.resolution
|
||||
while True:
|
||||
try:
|
||||
video, _, _, idx = self._clips.get_clip(idx)
|
||||
except Exception as e:
|
||||
print(idx, e)
|
||||
idx = (idx + 1) % self._clips.num_clips()
|
||||
continue
|
||||
break
|
||||
|
||||
return dict(**preprocess(video, resolution, sample_every_n_frames=self.sample_every_n_frames))
|
||||
|
||||
|
||||
class FrameDataset(data.Dataset):
|
||||
"""
|
||||
Generic dataset for videos stored as images. The loading will iterates over all the folders and subfolders
|
||||
in the provided directory. Each leaf folder is assumed to contain frames from a single video.
|
||||
|
||||
Args:
|
||||
data_folder: path to the folder with video frames. The folder
|
||||
should contain folders with frames from each video.
|
||||
sequence_length: length of extracted video sequences
|
||||
resolution: resolution of the returned videos
|
||||
sample_every_n_frames: sample every n frames from the video
|
||||
"""
|
||||
|
||||
def __init__(self, data_folder, sequence_length, resolution=64, sample_every_n_frames=1):
|
||||
self.resolution = resolution
|
||||
self.sequence_length = sequence_length
|
||||
self.sample_every_n_frames = sample_every_n_frames
|
||||
self.data_all = self.load_video_frames(data_folder)
|
||||
self.video_num = len(self.data_all)
|
||||
|
||||
def __getitem__(self, index):
|
||||
batch_data = self.getTensor(index)
|
||||
return_list = {'video': batch_data}
|
||||
|
||||
return return_list
|
||||
|
||||
def load_video_frames(self, dataroot: str) -> list:
|
||||
'''
|
||||
Loads all the video frames under the dataroot and returns a list of all the video frames.
|
||||
|
||||
Args:
|
||||
dataroot: The root directory containing the video frames.
|
||||
|
||||
Returns:
|
||||
A list of all the video frames.
|
||||
|
||||
'''
|
||||
data_all = []
|
||||
frame_list = os.walk(dataroot)
|
||||
for _, meta in enumerate(frame_list):
|
||||
root = meta[0]
|
||||
try:
|
||||
frames = sorted(meta[2], key=lambda item: int(item.split('.')[0].split('_')[-1]))
|
||||
except:
|
||||
print(meta[0], meta[2])
|
||||
if len(frames) < max(0, self.sequence_length * self.sample_every_n_frames):
|
||||
continue
|
||||
frames = [
|
||||
os.path.join(root, item) for item in frames
|
||||
if is_image_file(item)
|
||||
]
|
||||
if len(frames) > max(0, self.sequence_length * self.sample_every_n_frames):
|
||||
data_all.append(frames)
|
||||
|
||||
return data_all
|
||||
|
||||
def getTensor(self, index: int) -> torch.Tensor:
|
||||
'''
|
||||
Returns a tensor of the video frames at the given index.
|
||||
|
||||
Args:
|
||||
index: The index of the video frames to return.
|
||||
|
||||
Returns:
|
||||
A BCTHW tensor in the range `[0, 1]` of the video frames at the given index.
|
||||
|
||||
'''
|
||||
video = self.data_all[index]
|
||||
video_len = len(video)
|
||||
|
||||
# load the entire video when sequence_length = -1, whiel the sample_every_n_frames has to be 1
|
||||
if self.sequence_length == -1:
|
||||
assert self.sample_every_n_frames == 1
|
||||
start_idx = 0
|
||||
end_idx = video_len
|
||||
else:
|
||||
n_frames_interval = self.sequence_length * self.sample_every_n_frames
|
||||
start_idx = random.randint(0, video_len - n_frames_interval)
|
||||
end_idx = start_idx + n_frames_interval
|
||||
img = Image.open(video[0])
|
||||
h, w = img.height, img.width
|
||||
|
||||
if h > w:
|
||||
half = (h - w) // 2
|
||||
cropsize = (0, half, w, half + w) # left, upper, right, lower
|
||||
elif w > h:
|
||||
half = (w - h) // 2
|
||||
cropsize = (half, 0, half + h, h)
|
||||
|
||||
images = []
|
||||
for i in range(start_idx, end_idx,
|
||||
self.sample_every_n_frames):
|
||||
path = video[i]
|
||||
img = Image.open(path)
|
||||
|
||||
if h != w:
|
||||
img = img.crop(cropsize)
|
||||
|
||||
img = img.resize(
|
||||
(self.resolution, self.resolution),
|
||||
Image.ANTIALIAS)
|
||||
img = np.asarray(img, dtype=np.float32)
|
||||
img /= 255.
|
||||
img_tensor = preprocess_image(img).unsqueeze(0)
|
||||
images.append(img_tensor)
|
||||
|
||||
video_clip = torch.cat(images).permute(3, 0, 1, 2)
|
||||
return video_clip
|
||||
|
||||
def __len__(self):
|
||||
return self.video_num
|
||||
|
||||
@@ -1,161 +1,161 @@
|
||||
# Adapted from https://github.com/universome/stylegan-v/blob/master/src/metrics/metric_utils.py
|
||||
import os
|
||||
import random
|
||||
import torch
|
||||
import pickle
|
||||
import numpy as np
|
||||
|
||||
from typing import List, Tuple
|
||||
|
||||
def seed_everything(seed):
|
||||
random.seed(seed)
|
||||
os.environ['PYTHONHASHSEED'] = str(seed)
|
||||
np.random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed(seed)
|
||||
|
||||
|
||||
class FeatureStats:
|
||||
'''
|
||||
Class to store statistics of features, including all features and mean/covariance.
|
||||
|
||||
Args:
|
||||
capture_all: Whether to store all the features.
|
||||
capture_mean_cov: Whether to store mean and covariance.
|
||||
max_items: Maximum number of items to store.
|
||||
'''
|
||||
def __init__(self, capture_all: bool = False, capture_mean_cov: bool = False, max_items: int = None):
|
||||
'''
|
||||
'''
|
||||
self.capture_all = capture_all
|
||||
self.capture_mean_cov = capture_mean_cov
|
||||
self.max_items = max_items
|
||||
self.num_items = 0
|
||||
self.num_features = None
|
||||
self.all_features = None
|
||||
self.raw_mean = None
|
||||
self.raw_cov = None
|
||||
|
||||
def set_num_features(self, num_features: int):
|
||||
'''
|
||||
Set the number of features diminsions.
|
||||
|
||||
Args:
|
||||
num_features: Number of features diminsions.
|
||||
'''
|
||||
if self.num_features is not None:
|
||||
assert num_features == self.num_features
|
||||
else:
|
||||
self.num_features = num_features
|
||||
self.all_features = []
|
||||
self.raw_mean = np.zeros([num_features], dtype=np.float64)
|
||||
self.raw_cov = np.zeros([num_features, num_features], dtype=np.float64)
|
||||
|
||||
def is_full(self) -> bool:
|
||||
'''
|
||||
Check if the maximum number of samples is reached.
|
||||
|
||||
Returns:
|
||||
True if the storage is full, False otherwise.
|
||||
'''
|
||||
return (self.max_items is not None) and (self.num_items >= self.max_items)
|
||||
|
||||
def append(self, x: np.ndarray):
|
||||
'''
|
||||
Add the newly computed features to the list. Update the mean and covariance.
|
||||
|
||||
Args:
|
||||
x: New features to record.
|
||||
'''
|
||||
x = np.asarray(x, dtype=np.float32)
|
||||
assert x.ndim == 2
|
||||
if (self.max_items is not None) and (self.num_items + x.shape[0] > self.max_items):
|
||||
if self.num_items >= self.max_items:
|
||||
return
|
||||
x = x[:self.max_items - self.num_items]
|
||||
|
||||
self.set_num_features(x.shape[1])
|
||||
self.num_items += x.shape[0]
|
||||
if self.capture_all:
|
||||
self.all_features.append(x)
|
||||
if self.capture_mean_cov:
|
||||
x64 = x.astype(np.float64)
|
||||
self.raw_mean += x64.sum(axis=0)
|
||||
self.raw_cov += x64.T @ x64
|
||||
|
||||
def append_torch(self, x: torch.Tensor, rank: int, num_gpus: int):
|
||||
'''
|
||||
Add the newly computed PyTorch features to the list. Update the mean and covariance.
|
||||
|
||||
Args:
|
||||
x: New features to record.
|
||||
rank: Rank of the current GPU.
|
||||
num_gpus: Total number of GPUs.
|
||||
'''
|
||||
assert isinstance(x, torch.Tensor) and x.ndim == 2
|
||||
assert 0 <= rank < num_gpus
|
||||
if num_gpus > 1:
|
||||
ys = []
|
||||
for src in range(num_gpus):
|
||||
y = x.clone()
|
||||
torch.distributed.broadcast(y, src=src)
|
||||
ys.append(y)
|
||||
x = torch.stack(ys, dim=1).flatten(0, 1) # interleave samples
|
||||
self.append(x.cpu().numpy())
|
||||
|
||||
def get_all(self) -> np.ndarray:
|
||||
'''
|
||||
Get all the stored features as NumPy Array.
|
||||
|
||||
Returns:
|
||||
Concatenation of the stored features.
|
||||
'''
|
||||
assert self.capture_all
|
||||
return np.concatenate(self.all_features, axis=0)
|
||||
|
||||
def get_all_torch(self) -> torch.Tensor:
|
||||
'''
|
||||
Get all the stored features as PyTorch Tensor.
|
||||
|
||||
Returns:
|
||||
Concatenation of the stored features.
|
||||
'''
|
||||
return torch.from_numpy(self.get_all())
|
||||
|
||||
def get_mean_cov(self) -> Tuple[np.ndarray, np.ndarray]:
|
||||
'''
|
||||
Get the mean and covariance of the stored features.
|
||||
|
||||
Returns:
|
||||
Mean and covariance of the stored features.
|
||||
'''
|
||||
assert self.capture_mean_cov
|
||||
mean = self.raw_mean / self.num_items
|
||||
cov = self.raw_cov / self.num_items
|
||||
cov = cov - np.outer(mean, mean)
|
||||
return mean, cov
|
||||
|
||||
def save(self, pkl_file: str):
|
||||
'''
|
||||
Save the features and statistics to a pickle file.
|
||||
|
||||
Args:
|
||||
pkl_file: Path to the pickle file.
|
||||
'''
|
||||
with open(pkl_file, 'wb') as f:
|
||||
pickle.dump(self.__dict__, f)
|
||||
|
||||
@staticmethod
|
||||
def load(pkl_file: str) -> 'FeatureStats':
|
||||
'''
|
||||
Load the features and statistics from a pickle file.
|
||||
|
||||
Args:
|
||||
pkl_file: Path to the pickle file.
|
||||
'''
|
||||
with open(pkl_file, 'rb') as f:
|
||||
s = pickle.load(f)
|
||||
obj = FeatureStats(capture_all=s['capture_all'], max_items=s['max_items'])
|
||||
obj.__dict__.update(s)
|
||||
print('Loaded %d features from %s' % (obj.num_items, pkl_file))
|
||||
return obj
|
||||
# Adapted from https://github.com/universome/stylegan-v/blob/master/src/metrics/metric_utils.py
|
||||
import os
|
||||
import random
|
||||
import torch
|
||||
import pickle
|
||||
import numpy as np
|
||||
|
||||
from typing import List, Tuple
|
||||
|
||||
def seed_everything(seed):
|
||||
random.seed(seed)
|
||||
os.environ['PYTHONHASHSEED'] = str(seed)
|
||||
np.random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed(seed)
|
||||
|
||||
|
||||
class FeatureStats:
|
||||
'''
|
||||
Class to store statistics of features, including all features and mean/covariance.
|
||||
|
||||
Args:
|
||||
capture_all: Whether to store all the features.
|
||||
capture_mean_cov: Whether to store mean and covariance.
|
||||
max_items: Maximum number of items to store.
|
||||
'''
|
||||
def __init__(self, capture_all: bool = False, capture_mean_cov: bool = False, max_items: int = None):
|
||||
'''
|
||||
'''
|
||||
self.capture_all = capture_all
|
||||
self.capture_mean_cov = capture_mean_cov
|
||||
self.max_items = max_items
|
||||
self.num_items = 0
|
||||
self.num_features = None
|
||||
self.all_features = None
|
||||
self.raw_mean = None
|
||||
self.raw_cov = None
|
||||
|
||||
def set_num_features(self, num_features: int):
|
||||
'''
|
||||
Set the number of features diminsions.
|
||||
|
||||
Args:
|
||||
num_features: Number of features diminsions.
|
||||
'''
|
||||
if self.num_features is not None:
|
||||
assert num_features == self.num_features
|
||||
else:
|
||||
self.num_features = num_features
|
||||
self.all_features = []
|
||||
self.raw_mean = np.zeros([num_features], dtype=np.float64)
|
||||
self.raw_cov = np.zeros([num_features, num_features], dtype=np.float64)
|
||||
|
||||
def is_full(self) -> bool:
|
||||
'''
|
||||
Check if the maximum number of samples is reached.
|
||||
|
||||
Returns:
|
||||
True if the storage is full, False otherwise.
|
||||
'''
|
||||
return (self.max_items is not None) and (self.num_items >= self.max_items)
|
||||
|
||||
def append(self, x: np.ndarray):
|
||||
'''
|
||||
Add the newly computed features to the list. Update the mean and covariance.
|
||||
|
||||
Args:
|
||||
x: New features to record.
|
||||
'''
|
||||
x = np.asarray(x, dtype=np.float32)
|
||||
assert x.ndim == 2
|
||||
if (self.max_items is not None) and (self.num_items + x.shape[0] > self.max_items):
|
||||
if self.num_items >= self.max_items:
|
||||
return
|
||||
x = x[:self.max_items - self.num_items]
|
||||
|
||||
self.set_num_features(x.shape[1])
|
||||
self.num_items += x.shape[0]
|
||||
if self.capture_all:
|
||||
self.all_features.append(x)
|
||||
if self.capture_mean_cov:
|
||||
x64 = x.astype(np.float64)
|
||||
self.raw_mean += x64.sum(axis=0)
|
||||
self.raw_cov += x64.T @ x64
|
||||
|
||||
def append_torch(self, x: torch.Tensor, rank: int, num_gpus: int):
|
||||
'''
|
||||
Add the newly computed PyTorch features to the list. Update the mean and covariance.
|
||||
|
||||
Args:
|
||||
x: New features to record.
|
||||
rank: Rank of the current GPU.
|
||||
num_gpus: Total number of GPUs.
|
||||
'''
|
||||
assert isinstance(x, torch.Tensor) and x.ndim == 2
|
||||
assert 0 <= rank < num_gpus
|
||||
if num_gpus > 1:
|
||||
ys = []
|
||||
for src in range(num_gpus):
|
||||
y = x.clone()
|
||||
torch.distributed.broadcast(y, src=src)
|
||||
ys.append(y)
|
||||
x = torch.stack(ys, dim=1).flatten(0, 1) # interleave samples
|
||||
self.append(x.cpu().numpy())
|
||||
|
||||
def get_all(self) -> np.ndarray:
|
||||
'''
|
||||
Get all the stored features as NumPy Array.
|
||||
|
||||
Returns:
|
||||
Concatenation of the stored features.
|
||||
'''
|
||||
assert self.capture_all
|
||||
return np.concatenate(self.all_features, axis=0)
|
||||
|
||||
def get_all_torch(self) -> torch.Tensor:
|
||||
'''
|
||||
Get all the stored features as PyTorch Tensor.
|
||||
|
||||
Returns:
|
||||
Concatenation of the stored features.
|
||||
'''
|
||||
return torch.from_numpy(self.get_all())
|
||||
|
||||
def get_mean_cov(self) -> Tuple[np.ndarray, np.ndarray]:
|
||||
'''
|
||||
Get the mean and covariance of the stored features.
|
||||
|
||||
Returns:
|
||||
Mean and covariance of the stored features.
|
||||
'''
|
||||
assert self.capture_mean_cov
|
||||
mean = self.raw_mean / self.num_items
|
||||
cov = self.raw_cov / self.num_items
|
||||
cov = cov - np.outer(mean, mean)
|
||||
return mean, cov
|
||||
|
||||
def save(self, pkl_file: str):
|
||||
'''
|
||||
Save the features and statistics to a pickle file.
|
||||
|
||||
Args:
|
||||
pkl_file: Path to the pickle file.
|
||||
'''
|
||||
with open(pkl_file, 'wb') as f:
|
||||
pickle.dump(self.__dict__, f)
|
||||
|
||||
@staticmethod
|
||||
def load(pkl_file: str) -> 'FeatureStats':
|
||||
'''
|
||||
Load the features and statistics from a pickle file.
|
||||
|
||||
Args:
|
||||
pkl_file: Path to the pickle file.
|
||||
'''
|
||||
with open(pkl_file, 'rb') as f:
|
||||
s = pickle.load(f)
|
||||
obj = FeatureStats(capture_all=s['capture_all'], max_items=s['max_items'])
|
||||
obj.__dict__.update(s)
|
||||
print('Loaded %d features from %s' % (obj.num_items, pkl_file))
|
||||
return obj
|
||||
|
||||
@@ -1,138 +1,144 @@
|
||||
# Adapted from https://github.com/guanjz20/StyleSync/blob/main/utils.py
|
||||
|
||||
import numpy as np
|
||||
import cv2
|
||||
|
||||
|
||||
def transformation_from_points(points1, points0, smooth=True, p_bias=None):
|
||||
points2 = np.array(points0)
|
||||
points2 = points2.astype(np.float64)
|
||||
points1 = points1.astype(np.float64)
|
||||
c1 = np.mean(points1, axis=0)
|
||||
c2 = np.mean(points2, axis=0)
|
||||
points1 -= c1
|
||||
points2 -= c2
|
||||
s1 = np.std(points1)
|
||||
s2 = np.std(points2)
|
||||
points1 /= s1
|
||||
points2 /= s2
|
||||
U, S, Vt = np.linalg.svd(np.matmul(points1.T, points2))
|
||||
R = (np.matmul(U, Vt)).T
|
||||
sR = (s2 / s1) * R
|
||||
T = c2.reshape(2, 1) - (s2 / s1) * np.matmul(R, c1.reshape(2, 1))
|
||||
M = np.concatenate((sR, T), axis=1)
|
||||
if smooth:
|
||||
bias = points2[2] - points1[2]
|
||||
if p_bias is None:
|
||||
p_bias = bias
|
||||
else:
|
||||
bias = p_bias * 0.2 + bias * 0.8
|
||||
p_bias = bias
|
||||
M[:, 2] = M[:, 2] + bias
|
||||
return M, p_bias
|
||||
|
||||
|
||||
class AlignRestore(object):
|
||||
def __init__(self, align_points=3):
|
||||
if align_points == 3:
|
||||
self.upscale_factor = 1
|
||||
self.crop_ratio = (2.8, 2.8)
|
||||
self.face_template = np.array([[19 - 2, 30 - 10], [56 + 2, 30 - 10], [37.5, 45 - 5]])
|
||||
self.face_template = self.face_template * 2.8
|
||||
# self.face_size = (int(100 * self.crop_ratio[0]), int(100 * self.crop_ratio[1]))
|
||||
self.face_size = (int(75 * self.crop_ratio[0]), int(100 * self.crop_ratio[1]))
|
||||
self.p_bias = None
|
||||
|
||||
def process(self, img, lmk_align=None, smooth=True, align_points=3):
|
||||
aligned_face, affine_matrix = self.align_warp_face(img, lmk_align, smooth)
|
||||
restored_img = self.restore_img(img, aligned_face, affine_matrix)
|
||||
cv2.imwrite("restored.jpg", restored_img)
|
||||
cv2.imwrite("aligned.jpg", aligned_face)
|
||||
return aligned_face, restored_img
|
||||
|
||||
def align_warp_face(self, img, lmks3, smooth=True, border_mode="constant"):
|
||||
affine_matrix, self.p_bias = transformation_from_points(lmks3, self.face_template, smooth, self.p_bias)
|
||||
if border_mode == "constant":
|
||||
border_mode = cv2.BORDER_CONSTANT
|
||||
elif border_mode == "reflect101":
|
||||
border_mode = cv2.BORDER_REFLECT101
|
||||
elif border_mode == "reflect":
|
||||
border_mode = cv2.BORDER_REFLECT
|
||||
cropped_face = cv2.warpAffine(
|
||||
img, affine_matrix, self.face_size, borderMode=border_mode, borderValue=[127, 127, 127]
|
||||
)
|
||||
return cropped_face, affine_matrix
|
||||
|
||||
def align_warp_face2(self, img, landmark, border_mode="constant"):
|
||||
affine_matrix = cv2.estimateAffinePartial2D(landmark, self.face_template)[0]
|
||||
if border_mode == "constant":
|
||||
border_mode = cv2.BORDER_CONSTANT
|
||||
elif border_mode == "reflect101":
|
||||
border_mode = cv2.BORDER_REFLECT101
|
||||
elif border_mode == "reflect":
|
||||
border_mode = cv2.BORDER_REFLECT
|
||||
cropped_face = cv2.warpAffine(
|
||||
img, affine_matrix, self.face_size, borderMode=border_mode, borderValue=(135, 133, 132)
|
||||
)
|
||||
return cropped_face, affine_matrix
|
||||
|
||||
def restore_img(self, input_img, face, affine_matrix):
|
||||
h, w, _ = input_img.shape
|
||||
h_up, w_up = int(h * self.upscale_factor), int(w * self.upscale_factor)
|
||||
upsample_img = cv2.resize(input_img, (w_up, h_up), interpolation=cv2.INTER_LANCZOS4)
|
||||
inverse_affine = cv2.invertAffineTransform(affine_matrix)
|
||||
inverse_affine *= self.upscale_factor
|
||||
if self.upscale_factor > 1:
|
||||
extra_offset = 0.5 * self.upscale_factor
|
||||
else:
|
||||
extra_offset = 0
|
||||
inverse_affine[:, 2] += extra_offset
|
||||
inv_restored = cv2.warpAffine(face, inverse_affine, (w_up, h_up))
|
||||
mask = np.ones((self.face_size[1], self.face_size[0]), dtype=np.float32)
|
||||
inv_mask = cv2.warpAffine(mask, inverse_affine, (w_up, h_up))
|
||||
inv_mask_erosion = cv2.erode(
|
||||
inv_mask, np.ones((int(2 * self.upscale_factor), int(2 * self.upscale_factor)), np.uint8)
|
||||
)
|
||||
pasted_face = inv_mask_erosion[:, :, None] * inv_restored
|
||||
total_face_area = np.sum(inv_mask_erosion)
|
||||
w_edge = int(total_face_area**0.5) // 20
|
||||
erosion_radius = w_edge * 2
|
||||
inv_mask_center = cv2.erode(inv_mask_erosion, np.ones((erosion_radius, erosion_radius), np.uint8))
|
||||
blur_size = w_edge * 2
|
||||
inv_soft_mask = cv2.GaussianBlur(inv_mask_center, (blur_size + 1, blur_size + 1), 0)
|
||||
inv_soft_mask = inv_soft_mask[:, :, None]
|
||||
upsample_img = inv_soft_mask * pasted_face + (1 - inv_soft_mask) * upsample_img
|
||||
if np.max(upsample_img) > 256:
|
||||
upsample_img = upsample_img.astype(np.uint16)
|
||||
else:
|
||||
upsample_img = upsample_img.astype(np.uint8)
|
||||
return upsample_img
|
||||
|
||||
|
||||
class laplacianSmooth:
|
||||
def __init__(self, smoothAlpha=0.3):
|
||||
self.smoothAlpha = smoothAlpha
|
||||
self.pts_last = None
|
||||
|
||||
def smooth(self, pts_cur):
|
||||
if self.pts_last is None:
|
||||
self.pts_last = pts_cur.copy()
|
||||
return pts_cur.copy()
|
||||
x1 = min(pts_cur[:, 0])
|
||||
x2 = max(pts_cur[:, 0])
|
||||
y1 = min(pts_cur[:, 1])
|
||||
y2 = max(pts_cur[:, 1])
|
||||
width = x2 - x1
|
||||
pts_update = []
|
||||
for i in range(len(pts_cur)):
|
||||
x_new, y_new = pts_cur[i]
|
||||
x_old, y_old = self.pts_last[i]
|
||||
tmp = (x_new - x_old) ** 2 + (y_new - y_old) ** 2
|
||||
w = np.exp(-tmp / (width * self.smoothAlpha))
|
||||
x = x_old * w + x_new * (1 - w)
|
||||
y = y_old * w + y_new * (1 - w)
|
||||
pts_update.append([x, y])
|
||||
pts_update = np.array(pts_update)
|
||||
self.pts_last = pts_update.copy()
|
||||
|
||||
return pts_update
|
||||
# Adapted from https://github.com/guanjz20/StyleSync/blob/main/utils.py
|
||||
|
||||
import numpy as np
|
||||
import cv2
|
||||
|
||||
|
||||
def transformation_from_points(points1, points0, smooth=True, p_bias=None):
|
||||
points2 = np.array(points0)
|
||||
points2 = points2.astype(np.float64)
|
||||
points1 = points1.astype(np.float64)
|
||||
c1 = np.mean(points1, axis=0)
|
||||
c2 = np.mean(points2, axis=0)
|
||||
points1 -= c1
|
||||
points2 -= c2
|
||||
s1 = np.std(points1)
|
||||
s2 = np.std(points2)
|
||||
points1 /= s1
|
||||
points2 /= s2
|
||||
U, S, Vt = np.linalg.svd(np.matmul(points1.T, points2))
|
||||
R = (np.matmul(U, Vt)).T
|
||||
sR = (s2 / s1) * R
|
||||
T = c2.reshape(2, 1) - (s2 / s1) * np.matmul(R, c1.reshape(2, 1))
|
||||
M = np.concatenate((sR, T), axis=1)
|
||||
if smooth:
|
||||
bias = points2[2] - points1[2]
|
||||
if p_bias is None:
|
||||
p_bias = bias
|
||||
else:
|
||||
bias = p_bias * 0.2 + bias * 0.8
|
||||
p_bias = bias
|
||||
M[:, 2] = M[:, 2] + bias
|
||||
return M, p_bias
|
||||
|
||||
|
||||
class AlignRestore(object):
|
||||
def __init__(self, align_points=3):
|
||||
if align_points == 3:
|
||||
self.upscale_factor = 1
|
||||
ratio = 2.8
|
||||
self.crop_ratio = (ratio, ratio)
|
||||
self.face_template = np.array([[19 - 2, 30 - 10], [56 + 2, 30 - 10], [37.5, 45 - 5]])
|
||||
self.face_template = self.face_template * ratio
|
||||
self.face_size = (int(75 * self.crop_ratio[0]), int(100 * self.crop_ratio[1]))
|
||||
self.p_bias = None
|
||||
|
||||
def process(self, img, lmk_align=None, smooth=True, align_points=3):
|
||||
aligned_face, affine_matrix = self.align_warp_face(img, lmk_align, smooth)
|
||||
restored_img = self.restore_img(img, aligned_face, affine_matrix)
|
||||
cv2.imwrite("restored.jpg", restored_img)
|
||||
cv2.imwrite("aligned.jpg", aligned_face)
|
||||
return aligned_face, restored_img
|
||||
|
||||
def align_warp_face(self, img, lmks3, smooth=True, border_mode="constant"):
|
||||
affine_matrix, self.p_bias = transformation_from_points(lmks3, self.face_template, smooth, self.p_bias)
|
||||
if border_mode == "constant":
|
||||
border_mode = cv2.BORDER_CONSTANT
|
||||
elif border_mode == "reflect101":
|
||||
border_mode = cv2.BORDER_REFLECT101
|
||||
elif border_mode == "reflect":
|
||||
border_mode = cv2.BORDER_REFLECT
|
||||
|
||||
cropped_face = cv2.warpAffine(
|
||||
img,
|
||||
affine_matrix,
|
||||
self.face_size,
|
||||
flags=cv2.INTER_LANCZOS4,
|
||||
borderMode=border_mode,
|
||||
borderValue=[127, 127, 127],
|
||||
)
|
||||
return cropped_face, affine_matrix
|
||||
|
||||
def align_warp_face2(self, img, landmark, border_mode="constant"):
|
||||
affine_matrix = cv2.estimateAffinePartial2D(landmark, self.face_template)[0]
|
||||
if border_mode == "constant":
|
||||
border_mode = cv2.BORDER_CONSTANT
|
||||
elif border_mode == "reflect101":
|
||||
border_mode = cv2.BORDER_REFLECT101
|
||||
elif border_mode == "reflect":
|
||||
border_mode = cv2.BORDER_REFLECT
|
||||
cropped_face = cv2.warpAffine(
|
||||
img, affine_matrix, self.face_size, borderMode=border_mode, borderValue=(135, 133, 132)
|
||||
)
|
||||
return cropped_face, affine_matrix
|
||||
|
||||
def restore_img(self, input_img, face, affine_matrix):
|
||||
h, w, _ = input_img.shape
|
||||
h_up, w_up = int(h * self.upscale_factor), int(w * self.upscale_factor)
|
||||
upsample_img = cv2.resize(input_img, (w_up, h_up), interpolation=cv2.INTER_LANCZOS4)
|
||||
inverse_affine = cv2.invertAffineTransform(affine_matrix)
|
||||
inverse_affine *= self.upscale_factor
|
||||
if self.upscale_factor > 1:
|
||||
extra_offset = 0.5 * self.upscale_factor
|
||||
else:
|
||||
extra_offset = 0
|
||||
inverse_affine[:, 2] += extra_offset
|
||||
inv_restored = cv2.warpAffine(face, inverse_affine, (w_up, h_up), flags=cv2.INTER_LANCZOS4)
|
||||
mask = np.ones((self.face_size[1], self.face_size[0]), dtype=np.float32)
|
||||
inv_mask = cv2.warpAffine(mask, inverse_affine, (w_up, h_up))
|
||||
inv_mask_erosion = cv2.erode(
|
||||
inv_mask, np.ones((int(2 * self.upscale_factor), int(2 * self.upscale_factor)), np.uint8)
|
||||
)
|
||||
pasted_face = inv_mask_erosion[:, :, None] * inv_restored
|
||||
total_face_area = np.sum(inv_mask_erosion)
|
||||
w_edge = int(total_face_area**0.5) // 20
|
||||
erosion_radius = w_edge * 2
|
||||
inv_mask_center = cv2.erode(inv_mask_erosion, np.ones((erosion_radius, erosion_radius), np.uint8))
|
||||
blur_size = w_edge * 2
|
||||
inv_soft_mask = cv2.GaussianBlur(inv_mask_center, (blur_size + 1, blur_size + 1), 0)
|
||||
inv_soft_mask = inv_soft_mask[:, :, None]
|
||||
upsample_img = inv_soft_mask * pasted_face + (1 - inv_soft_mask) * upsample_img
|
||||
if np.max(upsample_img) > 256:
|
||||
upsample_img = upsample_img.astype(np.uint16)
|
||||
else:
|
||||
upsample_img = upsample_img.astype(np.uint8)
|
||||
return upsample_img
|
||||
|
||||
|
||||
class laplacianSmooth:
|
||||
def __init__(self, smoothAlpha=0.3):
|
||||
self.smoothAlpha = smoothAlpha
|
||||
self.pts_last = None
|
||||
|
||||
def smooth(self, pts_cur):
|
||||
if self.pts_last is None:
|
||||
self.pts_last = pts_cur.copy()
|
||||
return pts_cur.copy()
|
||||
x1 = min(pts_cur[:, 0])
|
||||
x2 = max(pts_cur[:, 0])
|
||||
y1 = min(pts_cur[:, 1])
|
||||
y2 = max(pts_cur[:, 1])
|
||||
width = x2 - x1
|
||||
pts_update = []
|
||||
for i in range(len(pts_cur)):
|
||||
x_new, y_new = pts_cur[i]
|
||||
x_old, y_old = self.pts_last[i]
|
||||
tmp = (x_new - x_old) ** 2 + (y_new - y_old) ** 2
|
||||
w = np.exp(-tmp / (width * self.smoothAlpha))
|
||||
x = x_old * w + x_new * (1 - w)
|
||||
y = y_old * w + y_new * (1 - w)
|
||||
pts_update.append([x, y])
|
||||
pts_update = np.array(pts_update)
|
||||
self.pts_last = pts_update.copy()
|
||||
|
||||
return pts_update
|
||||
|
||||
+194
-194
@@ -1,194 +1,194 @@
|
||||
# Adapted from https://github.com/Rudrabha/Wav2Lip/blob/master/audio.py
|
||||
|
||||
import librosa
|
||||
import librosa.filters
|
||||
import numpy as np
|
||||
from scipy import signal
|
||||
from scipy.io import wavfile
|
||||
from omegaconf import OmegaConf
|
||||
import torch
|
||||
|
||||
audio_config_path = "configs/audio.yaml"
|
||||
|
||||
config = OmegaConf.load(audio_config_path)
|
||||
|
||||
|
||||
def load_wav(path, sr):
|
||||
return librosa.core.load(path, sr=sr)[0]
|
||||
|
||||
|
||||
def save_wav(wav, path, sr):
|
||||
wav *= 32767 / max(0.01, np.max(np.abs(wav)))
|
||||
# proposed by @dsmiller
|
||||
wavfile.write(path, sr, wav.astype(np.int16))
|
||||
|
||||
|
||||
def save_wavenet_wav(wav, path, sr):
|
||||
librosa.output.write_wav(path, wav, sr=sr)
|
||||
|
||||
|
||||
def preemphasis(wav, k, preemphasize=True):
|
||||
if preemphasize:
|
||||
return signal.lfilter([1, -k], [1], wav)
|
||||
return wav
|
||||
|
||||
|
||||
def inv_preemphasis(wav, k, inv_preemphasize=True):
|
||||
if inv_preemphasize:
|
||||
return signal.lfilter([1], [1, -k], wav)
|
||||
return wav
|
||||
|
||||
|
||||
def get_hop_size():
|
||||
hop_size = config.audio.hop_size
|
||||
if hop_size is None:
|
||||
assert config.audio.frame_shift_ms is not None
|
||||
hop_size = int(config.audio.frame_shift_ms / 1000 * config.audio.sample_rate)
|
||||
return hop_size
|
||||
|
||||
|
||||
def linearspectrogram(wav):
|
||||
D = _stft(preemphasis(wav, config.audio.preemphasis, config.audio.preemphasize))
|
||||
S = _amp_to_db(np.abs(D)) - config.audio.ref_level_db
|
||||
|
||||
if config.audio.signal_normalization:
|
||||
return _normalize(S)
|
||||
return S
|
||||
|
||||
|
||||
def melspectrogram(wav):
|
||||
D = _stft(preemphasis(wav, config.audio.preemphasis, config.audio.preemphasize))
|
||||
S = _amp_to_db(_linear_to_mel(np.abs(D))) - config.audio.ref_level_db
|
||||
|
||||
if config.audio.signal_normalization:
|
||||
return _normalize(S)
|
||||
return S
|
||||
|
||||
|
||||
def _lws_processor():
|
||||
import lws
|
||||
|
||||
return lws.lws(config.audio.n_fft, get_hop_size(), fftsize=config.audio.win_size, mode="speech")
|
||||
|
||||
|
||||
def _stft(y):
|
||||
if config.audio.use_lws:
|
||||
return _lws_processor(config.audio).stft(y).T
|
||||
else:
|
||||
return librosa.stft(y=y, n_fft=config.audio.n_fft, hop_length=get_hop_size(), win_length=config.audio.win_size)
|
||||
|
||||
|
||||
##########################################################
|
||||
# Those are only correct when using lws!!! (This was messing with Wavenet quality for a long time!)
|
||||
def num_frames(length, fsize, fshift):
|
||||
"""Compute number of time frames of spectrogram"""
|
||||
pad = fsize - fshift
|
||||
if length % fshift == 0:
|
||||
M = (length + pad * 2 - fsize) // fshift + 1
|
||||
else:
|
||||
M = (length + pad * 2 - fsize) // fshift + 2
|
||||
return M
|
||||
|
||||
|
||||
def pad_lr(x, fsize, fshift):
|
||||
"""Compute left and right padding"""
|
||||
M = num_frames(len(x), fsize, fshift)
|
||||
pad = fsize - fshift
|
||||
T = len(x) + 2 * pad
|
||||
r = (M - 1) * fshift + fsize - T
|
||||
return pad, pad + r
|
||||
|
||||
|
||||
##########################################################
|
||||
# Librosa correct padding
|
||||
def librosa_pad_lr(x, fsize, fshift):
|
||||
return 0, (x.shape[0] // fshift + 1) * fshift - x.shape[0]
|
||||
|
||||
|
||||
# Conversions
|
||||
_mel_basis = None
|
||||
|
||||
|
||||
def _linear_to_mel(spectogram):
|
||||
global _mel_basis
|
||||
if _mel_basis is None:
|
||||
_mel_basis = _build_mel_basis()
|
||||
return np.dot(_mel_basis, spectogram)
|
||||
|
||||
|
||||
def _build_mel_basis():
|
||||
assert config.audio.fmax <= config.audio.sample_rate // 2
|
||||
return librosa.filters.mel(
|
||||
sr=config.audio.sample_rate,
|
||||
n_fft=config.audio.n_fft,
|
||||
n_mels=config.audio.num_mels,
|
||||
fmin=config.audio.fmin,
|
||||
fmax=config.audio.fmax,
|
||||
)
|
||||
|
||||
|
||||
def _amp_to_db(x):
|
||||
min_level = np.exp(config.audio.min_level_db / 20 * np.log(10))
|
||||
return 20 * np.log10(np.maximum(min_level, x))
|
||||
|
||||
|
||||
def _db_to_amp(x):
|
||||
return np.power(10.0, (x) * 0.05)
|
||||
|
||||
|
||||
def _normalize(S):
|
||||
if config.audio.allow_clipping_in_normalization:
|
||||
if config.audio.symmetric_mels:
|
||||
return np.clip(
|
||||
(2 * config.audio.max_abs_value) * ((S - config.audio.min_level_db) / (-config.audio.min_level_db))
|
||||
- config.audio.max_abs_value,
|
||||
-config.audio.max_abs_value,
|
||||
config.audio.max_abs_value,
|
||||
)
|
||||
else:
|
||||
return np.clip(
|
||||
config.audio.max_abs_value * ((S - config.audio.min_level_db) / (-config.audio.min_level_db)),
|
||||
0,
|
||||
config.audio.max_abs_value,
|
||||
)
|
||||
|
||||
assert S.max() <= 0 and S.min() - config.audio.min_level_db >= 0
|
||||
if config.audio.symmetric_mels:
|
||||
return (2 * config.audio.max_abs_value) * (
|
||||
(S - config.audio.min_level_db) / (-config.audio.min_level_db)
|
||||
) - config.audio.max_abs_value
|
||||
else:
|
||||
return config.audio.max_abs_value * ((S - config.audio.min_level_db) / (-config.audio.min_level_db))
|
||||
|
||||
|
||||
def _denormalize(D):
|
||||
if config.audio.allow_clipping_in_normalization:
|
||||
if config.audio.symmetric_mels:
|
||||
return (
|
||||
(np.clip(D, -config.audio.max_abs_value, config.audio.max_abs_value) + config.audio.max_abs_value)
|
||||
* -config.audio.min_level_db
|
||||
/ (2 * config.audio.max_abs_value)
|
||||
) + config.audio.min_level_db
|
||||
else:
|
||||
return (
|
||||
np.clip(D, 0, config.audio.max_abs_value) * -config.audio.min_level_db / config.audio.max_abs_value
|
||||
) + config.audio.min_level_db
|
||||
|
||||
if config.audio.symmetric_mels:
|
||||
return (
|
||||
(D + config.audio.max_abs_value) * -config.audio.min_level_db / (2 * config.audio.max_abs_value)
|
||||
) + config.audio.min_level_db
|
||||
else:
|
||||
return (D * -config.audio.min_level_db / config.audio.max_abs_value) + config.audio.min_level_db
|
||||
|
||||
|
||||
def get_melspec_overlap(audio_samples, melspec_length=52):
|
||||
mel_spec_overlap = melspectrogram(audio_samples.numpy())
|
||||
mel_spec_overlap = torch.from_numpy(mel_spec_overlap)
|
||||
i = 0
|
||||
mel_spec_overlap_list = []
|
||||
while i + melspec_length < mel_spec_overlap.shape[1] - 3:
|
||||
mel_spec_overlap_list.append(mel_spec_overlap[:, i : i + melspec_length].unsqueeze(0))
|
||||
i += 3
|
||||
mel_spec_overlap = torch.stack(mel_spec_overlap_list)
|
||||
return mel_spec_overlap
|
||||
# Adapted from https://github.com/Rudrabha/Wav2Lip/blob/master/audio.py
|
||||
|
||||
import librosa
|
||||
import librosa.filters
|
||||
import numpy as np
|
||||
from scipy import signal
|
||||
from scipy.io import wavfile
|
||||
from omegaconf import OmegaConf
|
||||
import torch
|
||||
|
||||
audio_config_path = "configs/audio.yaml"
|
||||
|
||||
config = OmegaConf.load(audio_config_path)
|
||||
|
||||
|
||||
def load_wav(path, sr):
|
||||
return librosa.core.load(path, sr=sr)[0]
|
||||
|
||||
|
||||
def save_wav(wav, path, sr):
|
||||
wav *= 32767 / max(0.01, np.max(np.abs(wav)))
|
||||
# proposed by @dsmiller
|
||||
wavfile.write(path, sr, wav.astype(np.int16))
|
||||
|
||||
|
||||
def save_wavenet_wav(wav, path, sr):
|
||||
librosa.output.write_wav(path, wav, sr=sr)
|
||||
|
||||
|
||||
def preemphasis(wav, k, preemphasize=True):
|
||||
if preemphasize:
|
||||
return signal.lfilter([1, -k], [1], wav)
|
||||
return wav
|
||||
|
||||
|
||||
def inv_preemphasis(wav, k, inv_preemphasize=True):
|
||||
if inv_preemphasize:
|
||||
return signal.lfilter([1], [1, -k], wav)
|
||||
return wav
|
||||
|
||||
|
||||
def get_hop_size():
|
||||
hop_size = config.audio.hop_size
|
||||
if hop_size is None:
|
||||
assert config.audio.frame_shift_ms is not None
|
||||
hop_size = int(config.audio.frame_shift_ms / 1000 * config.audio.sample_rate)
|
||||
return hop_size
|
||||
|
||||
|
||||
def linearspectrogram(wav):
|
||||
D = _stft(preemphasis(wav, config.audio.preemphasis, config.audio.preemphasize))
|
||||
S = _amp_to_db(np.abs(D)) - config.audio.ref_level_db
|
||||
|
||||
if config.audio.signal_normalization:
|
||||
return _normalize(S)
|
||||
return S
|
||||
|
||||
|
||||
def melspectrogram(wav):
|
||||
D = _stft(preemphasis(wav, config.audio.preemphasis, config.audio.preemphasize))
|
||||
S = _amp_to_db(_linear_to_mel(np.abs(D))) - config.audio.ref_level_db
|
||||
|
||||
if config.audio.signal_normalization:
|
||||
return _normalize(S)
|
||||
return S
|
||||
|
||||
|
||||
def _lws_processor():
|
||||
import lws
|
||||
|
||||
return lws.lws(config.audio.n_fft, get_hop_size(), fftsize=config.audio.win_size, mode="speech")
|
||||
|
||||
|
||||
def _stft(y):
|
||||
if config.audio.use_lws:
|
||||
return _lws_processor(config.audio).stft(y).T
|
||||
else:
|
||||
return librosa.stft(y=y, n_fft=config.audio.n_fft, hop_length=get_hop_size(), win_length=config.audio.win_size)
|
||||
|
||||
|
||||
##########################################################
|
||||
# Those are only correct when using lws!!! (This was messing with Wavenet quality for a long time!)
|
||||
def num_frames(length, fsize, fshift):
|
||||
"""Compute number of time frames of spectrogram"""
|
||||
pad = fsize - fshift
|
||||
if length % fshift == 0:
|
||||
M = (length + pad * 2 - fsize) // fshift + 1
|
||||
else:
|
||||
M = (length + pad * 2 - fsize) // fshift + 2
|
||||
return M
|
||||
|
||||
|
||||
def pad_lr(x, fsize, fshift):
|
||||
"""Compute left and right padding"""
|
||||
M = num_frames(len(x), fsize, fshift)
|
||||
pad = fsize - fshift
|
||||
T = len(x) + 2 * pad
|
||||
r = (M - 1) * fshift + fsize - T
|
||||
return pad, pad + r
|
||||
|
||||
|
||||
##########################################################
|
||||
# Librosa correct padding
|
||||
def librosa_pad_lr(x, fsize, fshift):
|
||||
return 0, (x.shape[0] // fshift + 1) * fshift - x.shape[0]
|
||||
|
||||
|
||||
# Conversions
|
||||
_mel_basis = None
|
||||
|
||||
|
||||
def _linear_to_mel(spectogram):
|
||||
global _mel_basis
|
||||
if _mel_basis is None:
|
||||
_mel_basis = _build_mel_basis()
|
||||
return np.dot(_mel_basis, spectogram)
|
||||
|
||||
|
||||
def _build_mel_basis():
|
||||
assert config.audio.fmax <= config.audio.sample_rate // 2
|
||||
return librosa.filters.mel(
|
||||
sr=config.audio.sample_rate,
|
||||
n_fft=config.audio.n_fft,
|
||||
n_mels=config.audio.num_mels,
|
||||
fmin=config.audio.fmin,
|
||||
fmax=config.audio.fmax,
|
||||
)
|
||||
|
||||
|
||||
def _amp_to_db(x):
|
||||
min_level = np.exp(config.audio.min_level_db / 20 * np.log(10))
|
||||
return 20 * np.log10(np.maximum(min_level, x))
|
||||
|
||||
|
||||
def _db_to_amp(x):
|
||||
return np.power(10.0, (x) * 0.05)
|
||||
|
||||
|
||||
def _normalize(S):
|
||||
if config.audio.allow_clipping_in_normalization:
|
||||
if config.audio.symmetric_mels:
|
||||
return np.clip(
|
||||
(2 * config.audio.max_abs_value) * ((S - config.audio.min_level_db) / (-config.audio.min_level_db))
|
||||
- config.audio.max_abs_value,
|
||||
-config.audio.max_abs_value,
|
||||
config.audio.max_abs_value,
|
||||
)
|
||||
else:
|
||||
return np.clip(
|
||||
config.audio.max_abs_value * ((S - config.audio.min_level_db) / (-config.audio.min_level_db)),
|
||||
0,
|
||||
config.audio.max_abs_value,
|
||||
)
|
||||
|
||||
assert S.max() <= 0 and S.min() - config.audio.min_level_db >= 0
|
||||
if config.audio.symmetric_mels:
|
||||
return (2 * config.audio.max_abs_value) * (
|
||||
(S - config.audio.min_level_db) / (-config.audio.min_level_db)
|
||||
) - config.audio.max_abs_value
|
||||
else:
|
||||
return config.audio.max_abs_value * ((S - config.audio.min_level_db) / (-config.audio.min_level_db))
|
||||
|
||||
|
||||
def _denormalize(D):
|
||||
if config.audio.allow_clipping_in_normalization:
|
||||
if config.audio.symmetric_mels:
|
||||
return (
|
||||
(np.clip(D, -config.audio.max_abs_value, config.audio.max_abs_value) + config.audio.max_abs_value)
|
||||
* -config.audio.min_level_db
|
||||
/ (2 * config.audio.max_abs_value)
|
||||
) + config.audio.min_level_db
|
||||
else:
|
||||
return (
|
||||
np.clip(D, 0, config.audio.max_abs_value) * -config.audio.min_level_db / config.audio.max_abs_value
|
||||
) + config.audio.min_level_db
|
||||
|
||||
if config.audio.symmetric_mels:
|
||||
return (
|
||||
(D + config.audio.max_abs_value) * -config.audio.min_level_db / (2 * config.audio.max_abs_value)
|
||||
) + config.audio.min_level_db
|
||||
else:
|
||||
return (D * -config.audio.min_level_db / config.audio.max_abs_value) + config.audio.min_level_db
|
||||
|
||||
|
||||
def get_melspec_overlap(audio_samples, melspec_length=52):
|
||||
mel_spec_overlap = melspectrogram(audio_samples.numpy())
|
||||
mel_spec_overlap = torch.from_numpy(mel_spec_overlap)
|
||||
i = 0
|
||||
mel_spec_overlap_list = []
|
||||
while i + melspec_length < mel_spec_overlap.shape[1] - 3:
|
||||
mel_spec_overlap_list.append(mel_spec_overlap[:, i : i + melspec_length].unsqueeze(0))
|
||||
i += 3
|
||||
mel_spec_overlap = torch.stack(mel_spec_overlap_list)
|
||||
return mel_spec_overlap
|
||||
|
||||
+157
-157
@@ -1,157 +1,157 @@
|
||||
# We modified the original AVReader class of decord to solve the problem of memory leak.
|
||||
# For more details, refer to: https://github.com/dmlc/decord/issues/208
|
||||
|
||||
import numpy as np
|
||||
from decord.video_reader import VideoReader
|
||||
from decord.audio_reader import AudioReader
|
||||
|
||||
from decord.ndarray import cpu
|
||||
from decord import ndarray as _nd
|
||||
from decord.bridge import bridge_out
|
||||
|
||||
|
||||
class AVReader(object):
|
||||
"""Individual audio video reader with convenient indexing function.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
uri: str
|
||||
Path of file.
|
||||
ctx: decord.Context
|
||||
The context to decode the file, can be decord.cpu() or decord.gpu().
|
||||
sample_rate: int, default is -1
|
||||
Desired output sample rate of the audio, unchanged if `-1` is specified.
|
||||
mono: bool, default is True
|
||||
Desired output channel layout of the audio. `True` is mono layout. `False` is unchanged.
|
||||
width : int, default is -1
|
||||
Desired output width of the video, unchanged if `-1` is specified.
|
||||
height : int, default is -1
|
||||
Desired output height of the video, unchanged if `-1` is specified.
|
||||
num_threads : int, default is 0
|
||||
Number of decoding thread, auto if `0` is specified.
|
||||
fault_tol : int, default is -1
|
||||
The threshold of corupted and recovered frames. This is to prevent silent fault
|
||||
tolerance when for example 50% frames of a video cannot be decoded and duplicate
|
||||
frames are returned. You may find the fault tolerant feature sweet in many cases,
|
||||
but not for training models. Say `N = # recovered frames`
|
||||
If `fault_tol` < 0, nothing will happen.
|
||||
If 0 < `fault_tol` < 1.0, if N > `fault_tol * len(video)`, raise `DECORDLimitReachedError`.
|
||||
If 1 < `fault_tol`, if N > `fault_tol`, raise `DECORDLimitReachedError`.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, uri, ctx=cpu(0), sample_rate=44100, mono=True, width=-1, height=-1, num_threads=0, fault_tol=-1
|
||||
):
|
||||
self.__audio_reader = AudioReader(uri, ctx, sample_rate, mono)
|
||||
self.__audio_reader.add_padding()
|
||||
if hasattr(uri, "read"):
|
||||
uri.seek(0)
|
||||
self.__video_reader = VideoReader(uri, ctx, width, height, num_threads, fault_tol)
|
||||
self.__video_reader.seek(0)
|
||||
|
||||
def __len__(self):
|
||||
"""Get length of the video. Note that sometimes FFMPEG reports inaccurate number of frames,
|
||||
we always follow what FFMPEG reports.
|
||||
Returns
|
||||
-------
|
||||
int
|
||||
The number of frames in the video file.
|
||||
"""
|
||||
return len(self.__video_reader)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
"""Get audio samples and video frame at `idx`.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
idx : int or slice
|
||||
The frame index, can be negative which means it will index backwards,
|
||||
or slice of frame indices.
|
||||
|
||||
Returns
|
||||
-------
|
||||
(ndarray/list of ndarray, ndarray)
|
||||
First element is samples of shape CxS or a list of length N containing samples of shape CxS,
|
||||
where N is the number of frames, C is the number of channels,
|
||||
S is the number of samples of the corresponding frame.
|
||||
|
||||
Second element is Frame of shape HxWx3 or batch of image frames with shape NxHxWx3,
|
||||
where N is the length of the slice.
|
||||
"""
|
||||
assert self.__video_reader is not None and self.__audio_reader is not None
|
||||
if isinstance(idx, slice):
|
||||
return self.get_batch(range(*idx.indices(len(self.__video_reader))))
|
||||
if idx < 0:
|
||||
idx += len(self.__video_reader)
|
||||
if idx >= len(self.__video_reader) or idx < 0:
|
||||
raise IndexError("Index: {} out of bound: {}".format(idx, len(self.__video_reader)))
|
||||
audio_start_idx, audio_end_idx = self.__video_reader.get_frame_timestamp(idx)
|
||||
audio_start_idx = self.__audio_reader._time_to_sample(audio_start_idx)
|
||||
audio_end_idx = self.__audio_reader._time_to_sample(audio_end_idx)
|
||||
results = (self.__audio_reader[audio_start_idx:audio_end_idx], self.__video_reader[idx])
|
||||
self.__video_reader.seek(0)
|
||||
return results
|
||||
|
||||
def get_batch(self, indices):
|
||||
"""Get entire batch of audio samples and video frames.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
indices : list of integers
|
||||
A list of frame indices. If negative indices detected, the indices will be indexed from backward
|
||||
Returns
|
||||
-------
|
||||
(list of ndarray, ndarray)
|
||||
First element is a list of length N containing samples of shape CxS,
|
||||
where N is the number of frames, C is the number of channels,
|
||||
S is the number of samples of the corresponding frame.
|
||||
|
||||
Second element is Frame of shape HxWx3 or batch of image frames with shape NxHxWx3,
|
||||
where N is the length of the slice.
|
||||
|
||||
"""
|
||||
assert self.__video_reader is not None and self.__audio_reader is not None
|
||||
indices = self._validate_indices(indices)
|
||||
audio_arr = []
|
||||
prev_video_idx = None
|
||||
prev_audio_end_idx = None
|
||||
for idx in list(indices):
|
||||
frame_start_time, frame_end_time = self.__video_reader.get_frame_timestamp(idx)
|
||||
# timestamp and sample conversion could have some error that could cause non-continuous audio
|
||||
# we detect if retrieving continuous frame and make the audio continuous
|
||||
if prev_video_idx and idx == prev_video_idx + 1:
|
||||
audio_start_idx = prev_audio_end_idx
|
||||
else:
|
||||
audio_start_idx = self.__audio_reader._time_to_sample(frame_start_time)
|
||||
audio_end_idx = self.__audio_reader._time_to_sample(frame_end_time)
|
||||
audio_arr.append(self.__audio_reader[audio_start_idx:audio_end_idx])
|
||||
prev_video_idx = idx
|
||||
prev_audio_end_idx = audio_end_idx
|
||||
results = (audio_arr, self.__video_reader.get_batch(indices))
|
||||
self.__video_reader.seek(0)
|
||||
return results
|
||||
|
||||
def _get_slice(self, sl):
|
||||
audio_arr = np.empty(shape=(self.__audio_reader.shape()[0], 0), dtype="float32")
|
||||
for idx in list(sl):
|
||||
audio_start_idx, audio_end_idx = self.__video_reader.get_frame_timestamp(idx)
|
||||
audio_start_idx = self.__audio_reader._time_to_sample(audio_start_idx)
|
||||
audio_end_idx = self.__audio_reader._time_to_sample(audio_end_idx)
|
||||
audio_arr = np.concatenate(
|
||||
(audio_arr, self.__audio_reader[audio_start_idx:audio_end_idx].asnumpy()), axis=1
|
||||
)
|
||||
results = (bridge_out(_nd.array(audio_arr)), self.__video_reader.get_batch(sl))
|
||||
self.__video_reader.seek(0)
|
||||
return results
|
||||
|
||||
def _validate_indices(self, indices):
|
||||
"""Validate int64 integers and convert negative integers to positive by backward search"""
|
||||
assert self.__video_reader is not None and self.__audio_reader is not None
|
||||
indices = np.array(indices, dtype=np.int64)
|
||||
# process negative indices
|
||||
indices[indices < 0] += len(self.__video_reader)
|
||||
if not (indices >= 0).all():
|
||||
raise IndexError("Invalid negative indices: {}".format(indices[indices < 0] + len(self.__video_reader)))
|
||||
if not (indices < len(self.__video_reader)).all():
|
||||
raise IndexError("Out of bound indices: {}".format(indices[indices >= len(self.__video_reader)]))
|
||||
return indices
|
||||
# We modified the original AVReader class of decord to solve the problem of memory leak.
|
||||
# For more details, refer to: https://github.com/dmlc/decord/issues/208
|
||||
|
||||
import numpy as np
|
||||
from decord.video_reader import VideoReader
|
||||
from decord.audio_reader import AudioReader
|
||||
|
||||
from decord.ndarray import cpu
|
||||
from decord import ndarray as _nd
|
||||
from decord.bridge import bridge_out
|
||||
|
||||
|
||||
class AVReader(object):
|
||||
"""Individual audio video reader with convenient indexing function.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
uri: str
|
||||
Path of file.
|
||||
ctx: decord.Context
|
||||
The context to decode the file, can be decord.cpu() or decord.gpu().
|
||||
sample_rate: int, default is -1
|
||||
Desired output sample rate of the audio, unchanged if `-1` is specified.
|
||||
mono: bool, default is True
|
||||
Desired output channel layout of the audio. `True` is mono layout. `False` is unchanged.
|
||||
width : int, default is -1
|
||||
Desired output width of the video, unchanged if `-1` is specified.
|
||||
height : int, default is -1
|
||||
Desired output height of the video, unchanged if `-1` is specified.
|
||||
num_threads : int, default is 0
|
||||
Number of decoding thread, auto if `0` is specified.
|
||||
fault_tol : int, default is -1
|
||||
The threshold of corupted and recovered frames. This is to prevent silent fault
|
||||
tolerance when for example 50% frames of a video cannot be decoded and duplicate
|
||||
frames are returned. You may find the fault tolerant feature sweet in many cases,
|
||||
but not for training models. Say `N = # recovered frames`
|
||||
If `fault_tol` < 0, nothing will happen.
|
||||
If 0 < `fault_tol` < 1.0, if N > `fault_tol * len(video)`, raise `DECORDLimitReachedError`.
|
||||
If 1 < `fault_tol`, if N > `fault_tol`, raise `DECORDLimitReachedError`.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, uri, ctx=cpu(0), sample_rate=44100, mono=True, width=-1, height=-1, num_threads=0, fault_tol=-1
|
||||
):
|
||||
self.__audio_reader = AudioReader(uri, ctx, sample_rate, mono)
|
||||
self.__audio_reader.add_padding()
|
||||
if hasattr(uri, "read"):
|
||||
uri.seek(0)
|
||||
self.__video_reader = VideoReader(uri, ctx, width, height, num_threads, fault_tol)
|
||||
self.__video_reader.seek(0)
|
||||
|
||||
def __len__(self):
|
||||
"""Get length of the video. Note that sometimes FFMPEG reports inaccurate number of frames,
|
||||
we always follow what FFMPEG reports.
|
||||
Returns
|
||||
-------
|
||||
int
|
||||
The number of frames in the video file.
|
||||
"""
|
||||
return len(self.__video_reader)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
"""Get audio samples and video frame at `idx`.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
idx : int or slice
|
||||
The frame index, can be negative which means it will index backwards,
|
||||
or slice of frame indices.
|
||||
|
||||
Returns
|
||||
-------
|
||||
(ndarray/list of ndarray, ndarray)
|
||||
First element is samples of shape CxS or a list of length N containing samples of shape CxS,
|
||||
where N is the number of frames, C is the number of channels,
|
||||
S is the number of samples of the corresponding frame.
|
||||
|
||||
Second element is Frame of shape HxWx3 or batch of image frames with shape NxHxWx3,
|
||||
where N is the length of the slice.
|
||||
"""
|
||||
assert self.__video_reader is not None and self.__audio_reader is not None
|
||||
if isinstance(idx, slice):
|
||||
return self.get_batch(range(*idx.indices(len(self.__video_reader))))
|
||||
if idx < 0:
|
||||
idx += len(self.__video_reader)
|
||||
if idx >= len(self.__video_reader) or idx < 0:
|
||||
raise IndexError("Index: {} out of bound: {}".format(idx, len(self.__video_reader)))
|
||||
audio_start_idx, audio_end_idx = self.__video_reader.get_frame_timestamp(idx)
|
||||
audio_start_idx = self.__audio_reader._time_to_sample(audio_start_idx)
|
||||
audio_end_idx = self.__audio_reader._time_to_sample(audio_end_idx)
|
||||
results = (self.__audio_reader[audio_start_idx:audio_end_idx], self.__video_reader[idx])
|
||||
self.__video_reader.seek(0)
|
||||
return results
|
||||
|
||||
def get_batch(self, indices):
|
||||
"""Get entire batch of audio samples and video frames.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
indices : list of integers
|
||||
A list of frame indices. If negative indices detected, the indices will be indexed from backward
|
||||
Returns
|
||||
-------
|
||||
(list of ndarray, ndarray)
|
||||
First element is a list of length N containing samples of shape CxS,
|
||||
where N is the number of frames, C is the number of channels,
|
||||
S is the number of samples of the corresponding frame.
|
||||
|
||||
Second element is Frame of shape HxWx3 or batch of image frames with shape NxHxWx3,
|
||||
where N is the length of the slice.
|
||||
|
||||
"""
|
||||
assert self.__video_reader is not None and self.__audio_reader is not None
|
||||
indices = self._validate_indices(indices)
|
||||
audio_arr = []
|
||||
prev_video_idx = None
|
||||
prev_audio_end_idx = None
|
||||
for idx in list(indices):
|
||||
frame_start_time, frame_end_time = self.__video_reader.get_frame_timestamp(idx)
|
||||
# timestamp and sample conversion could have some error that could cause non-continuous audio
|
||||
# we detect if retrieving continuous frame and make the audio continuous
|
||||
if prev_video_idx and idx == prev_video_idx + 1:
|
||||
audio_start_idx = prev_audio_end_idx
|
||||
else:
|
||||
audio_start_idx = self.__audio_reader._time_to_sample(frame_start_time)
|
||||
audio_end_idx = self.__audio_reader._time_to_sample(frame_end_time)
|
||||
audio_arr.append(self.__audio_reader[audio_start_idx:audio_end_idx])
|
||||
prev_video_idx = idx
|
||||
prev_audio_end_idx = audio_end_idx
|
||||
results = (audio_arr, self.__video_reader.get_batch(indices))
|
||||
self.__video_reader.seek(0)
|
||||
return results
|
||||
|
||||
def _get_slice(self, sl):
|
||||
audio_arr = np.empty(shape=(self.__audio_reader.shape()[0], 0), dtype="float32")
|
||||
for idx in list(sl):
|
||||
audio_start_idx, audio_end_idx = self.__video_reader.get_frame_timestamp(idx)
|
||||
audio_start_idx = self.__audio_reader._time_to_sample(audio_start_idx)
|
||||
audio_end_idx = self.__audio_reader._time_to_sample(audio_end_idx)
|
||||
audio_arr = np.concatenate(
|
||||
(audio_arr, self.__audio_reader[audio_start_idx:audio_end_idx].asnumpy()), axis=1
|
||||
)
|
||||
results = (bridge_out(_nd.array(audio_arr)), self.__video_reader.get_batch(sl))
|
||||
self.__video_reader.seek(0)
|
||||
return results
|
||||
|
||||
def _validate_indices(self, indices):
|
||||
"""Validate int64 integers and convert negative integers to positive by backward search"""
|
||||
assert self.__video_reader is not None and self.__audio_reader is not None
|
||||
indices = np.array(indices, dtype=np.int64)
|
||||
# process negative indices
|
||||
indices[indices < 0] += len(self.__video_reader)
|
||||
if not (indices >= 0).all():
|
||||
raise IndexError("Invalid negative indices: {}".format(indices[indices < 0] + len(self.__video_reader)))
|
||||
if not (indices < len(self.__video_reader)).all():
|
||||
raise IndexError("Out of bound indices: {}".format(indices[indices >= len(self.__video_reader)]))
|
||||
return indices
|
||||
|
||||
+344
-241
@@ -1,241 +1,344 @@
|
||||
# Copyright (c) 2024 Bytedance Ltd. and/or its affiliates
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import os
|
||||
from torchvision import transforms
|
||||
import cv2
|
||||
from einops import rearrange
|
||||
import mediapipe as mp
|
||||
import torch
|
||||
import numpy as np
|
||||
from typing import Union
|
||||
from .affine_transform import AlignRestore, laplacianSmooth
|
||||
import face_alignment
|
||||
|
||||
"""
|
||||
If you are enlarging the image, you should prefer to use INTER_LINEAR or INTER_CUBIC interpolation. If you are shrinking the image, you should prefer to use INTER_AREA interpolation.
|
||||
https://stackoverflow.com/questions/23853632/which-kind-of-interpolation-best-for-resizing-image
|
||||
"""
|
||||
|
||||
def load_fixed_mask(resolution: int) -> torch.Tensor:
|
||||
mask_image = cv2.imread(os.path.join(os.path.dirname(__file__), "mask.png"))
|
||||
mask_image = cv2.cvtColor(mask_image, cv2.COLOR_BGR2RGB)
|
||||
mask_image = cv2.resize(mask_image, (resolution, resolution), interpolation=cv2.INTER_AREA) / 255.0
|
||||
mask_image = rearrange(torch.from_numpy(mask_image), "h w c -> c h w")
|
||||
return mask_image
|
||||
|
||||
class ImageProcessor:
|
||||
def __init__(self, resolution: int = 512, mask: str = "fix_mask", device: str = "cpu", mask_image=None):
|
||||
self.resolution = resolution
|
||||
self.resize = transforms.Resize(
|
||||
(resolution, resolution), interpolation=transforms.InterpolationMode.BILINEAR, antialias=True
|
||||
)
|
||||
self.normalize = transforms.Normalize([0.5], [0.5], inplace=True)
|
||||
self.mask = mask
|
||||
|
||||
if mask in ["mouth", "face", "eye"]:
|
||||
self.face_mesh = mp.solutions.face_mesh.FaceMesh(static_image_mode=True) # Process single image
|
||||
if mask == "fix_mask":
|
||||
self.face_mesh = None
|
||||
self.smoother = laplacianSmooth()
|
||||
self.restorer = AlignRestore()
|
||||
|
||||
if mask_image is None:
|
||||
self.mask_image = load_fixed_mask(resolution)
|
||||
else:
|
||||
self.mask_image = mask_image
|
||||
|
||||
if device != "cpu":
|
||||
self.fa = face_alignment.FaceAlignment(
|
||||
face_alignment.LandmarksType.TWO_D, flip_input=False, device=device
|
||||
)
|
||||
self.face_mesh = None
|
||||
else:
|
||||
self.face_mesh = None
|
||||
self.fa = None
|
||||
|
||||
def detect_facial_landmarks(self, image: np.ndarray):
|
||||
height, width, _ = image.shape
|
||||
results = self.face_mesh.process(image)
|
||||
if not results.multi_face_landmarks: # Face not detected
|
||||
print("Skipping frame: No face detected")
|
||||
return None # Return None instead of raising an error
|
||||
face_landmarks = results.multi_face_landmarks[0] # Only use the first face in the image
|
||||
landmark_coordinates = [
|
||||
(int(landmark.x * width), int(landmark.y * height)) for landmark in face_landmarks.landmark
|
||||
] # x means width, y means height
|
||||
return landmark_coordinates
|
||||
|
||||
def preprocess_one_masked_image(self, image: torch.Tensor) -> np.ndarray:
|
||||
image = self.resize(image)
|
||||
|
||||
if self.mask == "mouth" or self.mask == "face":
|
||||
landmark_coordinates = self.detect_facial_landmarks(image)
|
||||
if landmark_coordinates is None: # No face detected
|
||||
return None, None, None # Skip this frame
|
||||
|
||||
if self.mask == "mouth":
|
||||
surround_landmarks = mouth_surround_landmarks
|
||||
else:
|
||||
surround_landmarks = face_surround_landmarks
|
||||
|
||||
points = [landmark_coordinates[landmark] for landmark in surround_landmarks]
|
||||
points = np.array(points)
|
||||
mask = np.ones((self.resolution, self.resolution))
|
||||
mask = cv2.fillPoly(mask, pts=[points], color=(0, 0, 0))
|
||||
mask = torch.from_numpy(mask)
|
||||
mask = mask.unsqueeze(0)
|
||||
elif self.mask == "half":
|
||||
mask = torch.ones((self.resolution, self.resolution))
|
||||
height = mask.shape[0]
|
||||
mask[height // 2 :, :] = 0
|
||||
mask = mask.unsqueeze(0)
|
||||
elif self.mask == "eye":
|
||||
mask = torch.ones((self.resolution, self.resolution))
|
||||
landmark_coordinates = self.detect_facial_landmarks(image)
|
||||
if landmark_coordinates is None: # No face detected
|
||||
return None, None, None # Skip this frame
|
||||
y = landmark_coordinates[195][1]
|
||||
mask[y:, :] = 0
|
||||
mask = mask.unsqueeze(0)
|
||||
else:
|
||||
raise ValueError("Invalid mask type")
|
||||
|
||||
image = image.to(dtype=torch.float32)
|
||||
pixel_values = self.normalize(image / 255.0)
|
||||
masked_pixel_values = pixel_values * mask
|
||||
mask = 1 - mask
|
||||
|
||||
return pixel_values, masked_pixel_values, mask
|
||||
|
||||
def affine_transform(self, image: torch.Tensor):
|
||||
# Convert image to numpy array if necessary
|
||||
if isinstance(image, torch.Tensor):
|
||||
image = rearrange(image, "c h w -> h w c").numpy()
|
||||
|
||||
# Detect facial landmarks
|
||||
if self.fa is None:
|
||||
landmark_coordinates = self.detect_facial_landmarks(image)
|
||||
if landmark_coordinates is None: # No face detected
|
||||
return None, None, None # Skip this frame
|
||||
lm68 = mediapipe_lm478_to_face_alignment_lm68(landmark_coordinates)
|
||||
else:
|
||||
detected_faces = self.fa.get_landmarks(image)
|
||||
if detected_faces is None: # No face detected
|
||||
return None, None, None # Skip this frame
|
||||
lm68 = detected_faces[0]
|
||||
|
||||
# Perform affine transformation
|
||||
points = self.smoother.smooth(lm68)
|
||||
lmk3_ = np.zeros((3, 2))
|
||||
lmk3_[0] = points[17:22].mean(0)
|
||||
lmk3_[1] = points[22:27].mean(0)
|
||||
lmk3_[2] = points[27:36].mean(0)
|
||||
face, affine_matrix = self.restorer.align_warp_face(
|
||||
image.copy(), lmks3=lmk3_, smooth=True, border_mode="constant"
|
||||
)
|
||||
box = [0, 0, face.shape[1], face.shape[0]] # x1, y1, x2, y2
|
||||
face = cv2.resize(face, (self.resolution, self.resolution), interpolation=cv2.INTER_CUBIC)
|
||||
face = rearrange(torch.from_numpy(face), "h w c -> c h w")
|
||||
return face, box, affine_matrix
|
||||
|
||||
def preprocess_fixed_mask_image(self, image: torch.Tensor, affine_transform=False):
|
||||
if affine_transform:
|
||||
result = self.affine_transform(image)
|
||||
if result is None: # No face detected
|
||||
return None, None, None # Skip this frame
|
||||
image, _, _ = result
|
||||
else:
|
||||
image = self.resize(image)
|
||||
pixel_values = self.normalize(image / 255.0)
|
||||
masked_pixel_values = pixel_values * self.mask_image
|
||||
return pixel_values, masked_pixel_values, self.mask_image[0:1]
|
||||
|
||||
def prepare_masks_and_masked_images(self, images: Union[torch.Tensor, np.ndarray], affine_transform=False):
|
||||
if isinstance(images, np.ndarray):
|
||||
images = torch.from_numpy(images)
|
||||
if images.shape[3] == 3:
|
||||
images = rearrange(images, "b h w c -> b c h w")
|
||||
|
||||
pixel_values_list, masked_pixel_values_list, masks_list = [], [], []
|
||||
for image in images:
|
||||
if self.mask == "fix_mask":
|
||||
result = self.preprocess_fixed_mask_image(image, affine_transform=affine_transform)
|
||||
else:
|
||||
result = self.preprocess_one_masked_image(image)
|
||||
|
||||
if result is not None: # Skip frames where no face is detected
|
||||
pixel_values, masked_pixel_values, mask = result
|
||||
pixel_values_list.append(pixel_values)
|
||||
masked_pixel_values_list.append(masked_pixel_values)
|
||||
masks_list.append(mask)
|
||||
|
||||
if not pixel_values_list: # If no valid frames were processed
|
||||
return None, None, None
|
||||
|
||||
return torch.stack(pixel_values_list), torch.stack(masked_pixel_values_list), torch.stack(masks_list)
|
||||
|
||||
def process_images(self, images: Union[torch.Tensor, np.ndarray]):
|
||||
if isinstance(images, np.ndarray):
|
||||
images = torch.from_numpy(images)
|
||||
if images.shape[3] == 3:
|
||||
images = rearrange(images, "b h w c -> b c h w")
|
||||
images = self.resize(images)
|
||||
pixel_values = self.normalize(images / 255.0)
|
||||
return pixel_values
|
||||
|
||||
def close(self):
|
||||
if self.face_mesh is not None:
|
||||
self.face_mesh.close()
|
||||
|
||||
def mediapipe_lm478_to_face_alignment_lm68(lm478, return_2d=True):
|
||||
"""
|
||||
lm478: [B, 478, 3] or [478,3]
|
||||
"""
|
||||
landmarks_extracted = []
|
||||
for index in landmark_points_68:
|
||||
x = lm478[index][0]
|
||||
y = lm478[index][1]
|
||||
landmarks_extracted.append((x, y))
|
||||
return np.array(landmarks_extracted)
|
||||
|
||||
landmark_points_68 = [
|
||||
162, 234, 93, 58, 172, 136, 149, 148, 152, 377, 378, 365, 397, 288, 323, 454, 389, 71, 63, 105, 66, 107, 336, 296, 334, 293, 301, 168, 197, 5, 4, 75, 97, 2, 326, 305, 33, 160, 158, 133, 153, 144, 362, 385, 387, 263, 373, 380, 61, 39, 37, 0, 267, 269, 291, 405, 314, 17, 84, 181, 78, 82, 13, 312, 308, 317, 14, 87,
|
||||
]
|
||||
|
||||
# Refer to https://storage.googleapis.com/mediapipe-assets/documentation/mediapipe_face_landmark_fullsize.png
|
||||
mouth_surround_landmarks = [
|
||||
164, 165, 167, 92, 186, 57, 43, 106, 182, 83, 18, 313, 406, 335, 273, 287, 410, 322, 391, 393,
|
||||
]
|
||||
|
||||
face_surround_landmarks = [
|
||||
152, 377, 400, 378, 379, 365, 397, 288, 435, 433, 411, 425, 423, 327, 326, 94, 97, 98, 203, 205, 187, 213, 215, 58, 172, 136, 150, 149, 176, 148,
|
||||
]
|
||||
|
||||
if __name__ == "__main__":
|
||||
image_processor = ImageProcessor(512, mask="fix_mask")
|
||||
video = cv2.VideoCapture("/mnt/bn/maliva-gen-ai-v2/chunyu.li/HDTF/original/val/RD_Radio57_000.mp4")
|
||||
while True:
|
||||
ret, frame = video.read()
|
||||
if not ret:
|
||||
break
|
||||
|
||||
frame = rearrange(torch.Tensor(frame).type(torch.uint8), "h w c -> c h w")
|
||||
result = image_processor.affine_transform(frame)
|
||||
|
||||
if result is not None: # Only process frames where a face is detected
|
||||
face, _, _ = result
|
||||
face = (rearrange(face, "c h w -> h w c").detach().cpu().numpy()).astype(np.uint8)
|
||||
cv2.imwrite("face.jpg", face)
|
||||
break
|
||||
# Copyright (c) 2024 Bytedance Ltd. and/or its affiliates
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from torchvision import transforms
|
||||
import cv2
|
||||
from einops import rearrange
|
||||
import mediapipe as mp
|
||||
import torch
|
||||
import numpy as np
|
||||
from typing import Union
|
||||
from .affine_transform import AlignRestore, laplacianSmooth
|
||||
import face_alignment
|
||||
|
||||
"""
|
||||
If you are enlarging the image, you should prefer to use INTER_LINEAR or INTER_CUBIC interpolation. If you are shrinking the image, you should prefer to use INTER_AREA interpolation.
|
||||
https://stackoverflow.com/questions/23853632/which-kind-of-interpolation-best-for-resizing-image
|
||||
"""
|
||||
|
||||
|
||||
def load_fixed_mask(resolution: int, mask_image_path="latentsync/utils/mask.png") -> torch.Tensor:
|
||||
mask_image = cv2.imread(mask_image_path)
|
||||
mask_image = cv2.cvtColor(mask_image, cv2.COLOR_BGR2RGB)
|
||||
mask_image = cv2.resize(mask_image, (resolution, resolution), interpolation=cv2.INTER_LANCZOS4) / 255.0
|
||||
mask_image = rearrange(torch.from_numpy(mask_image), "h w c -> c h w")
|
||||
return mask_image
|
||||
|
||||
|
||||
class ImageProcessor:
|
||||
def __init__(self, resolution: int = 512, mask: str = "fix_mask", device: str = "cpu", mask_image=None):
|
||||
self.resolution = resolution
|
||||
self.resize = transforms.Resize(
|
||||
(resolution, resolution), interpolation=transforms.InterpolationMode.BILINEAR, antialias=True
|
||||
)
|
||||
self.normalize = transforms.Normalize([0.5], [0.5], inplace=True)
|
||||
self.mask = mask
|
||||
|
||||
if mask in ["mouth", "face", "eye"]:
|
||||
self.face_mesh = mp.solutions.face_mesh.FaceMesh(static_image_mode=True) # Process single image
|
||||
if mask == "fix_mask":
|
||||
self.face_mesh = None
|
||||
self.smoother = laplacianSmooth()
|
||||
self.restorer = AlignRestore()
|
||||
|
||||
if mask_image is None:
|
||||
self.mask_image = load_fixed_mask(resolution)
|
||||
else:
|
||||
self.mask_image = mask_image
|
||||
|
||||
if device != "cpu":
|
||||
self.fa = face_alignment.FaceAlignment(
|
||||
face_alignment.LandmarksType.TWO_D, flip_input=False, device=device
|
||||
)
|
||||
self.face_mesh = None
|
||||
else:
|
||||
# self.face_mesh = mp.solutions.face_mesh.FaceMesh(static_image_mode=True) # Process single image
|
||||
self.face_mesh = None
|
||||
self.fa = None
|
||||
|
||||
def detect_facial_landmarks(self, image: np.ndarray):
|
||||
height, width, _ = image.shape
|
||||
results = self.face_mesh.process(image)
|
||||
if not results.multi_face_landmarks: # Face not detected
|
||||
raise RuntimeError("Face not detected")
|
||||
face_landmarks = results.multi_face_landmarks[0] # Only use the first face in the image
|
||||
landmark_coordinates = [
|
||||
(int(landmark.x * width), int(landmark.y * height)) for landmark in face_landmarks.landmark
|
||||
] # x means width, y means height
|
||||
return landmark_coordinates
|
||||
|
||||
def preprocess_one_masked_image(self, image: torch.Tensor) -> np.ndarray:
|
||||
image = self.resize(image)
|
||||
|
||||
if self.mask == "mouth" or self.mask == "face":
|
||||
landmark_coordinates = self.detect_facial_landmarks(image)
|
||||
if self.mask == "mouth":
|
||||
surround_landmarks = mouth_surround_landmarks
|
||||
else:
|
||||
surround_landmarks = face_surround_landmarks
|
||||
|
||||
points = [landmark_coordinates[landmark] for landmark in surround_landmarks]
|
||||
points = np.array(points)
|
||||
mask = np.ones((self.resolution, self.resolution))
|
||||
mask = cv2.fillPoly(mask, pts=[points], color=(0, 0, 0))
|
||||
mask = torch.from_numpy(mask)
|
||||
mask = mask.unsqueeze(0)
|
||||
elif self.mask == "half":
|
||||
mask = torch.ones((self.resolution, self.resolution))
|
||||
height = mask.shape[0]
|
||||
mask[height // 2 :, :] = 0
|
||||
mask = mask.unsqueeze(0)
|
||||
elif self.mask == "eye":
|
||||
mask = torch.ones((self.resolution, self.resolution))
|
||||
landmark_coordinates = self.detect_facial_landmarks(image)
|
||||
y = landmark_coordinates[195][1]
|
||||
mask[y:, :] = 0
|
||||
mask = mask.unsqueeze(0)
|
||||
else:
|
||||
raise ValueError("Invalid mask type")
|
||||
|
||||
image = image.to(dtype=torch.float32)
|
||||
pixel_values = self.normalize(image / 255.0)
|
||||
masked_pixel_values = pixel_values * mask
|
||||
mask = 1 - mask
|
||||
|
||||
return pixel_values, masked_pixel_values, mask
|
||||
|
||||
def affine_transform(self, image: torch.Tensor, allow_multi_faces: bool = True) -> np.ndarray:
|
||||
# image = rearrange(image, "c h w-> h w c").numpy()
|
||||
if self.fa is None:
|
||||
landmark_coordinates = np.array(self.detect_facial_landmarks(image))
|
||||
lm68 = mediapipe_lm478_to_face_alignment_lm68(landmark_coordinates)
|
||||
else:
|
||||
detected_faces = self.fa.get_landmarks(image)
|
||||
if detected_faces is None:
|
||||
raise RuntimeError("Face not detected")
|
||||
if not allow_multi_faces and len(detected_faces) > 1:
|
||||
raise RuntimeError("More than one face detected")
|
||||
lm68 = detected_faces[0]
|
||||
|
||||
points = self.smoother.smooth(lm68)
|
||||
lmk3_ = np.zeros((3, 2))
|
||||
lmk3_[0] = points[17:22].mean(0)
|
||||
lmk3_[1] = points[22:27].mean(0)
|
||||
lmk3_[2] = points[27:36].mean(0)
|
||||
# print(lmk3_)
|
||||
face, affine_matrix = self.restorer.align_warp_face(
|
||||
image.copy(), lmks3=lmk3_, smooth=True, border_mode="constant"
|
||||
)
|
||||
box = [0, 0, face.shape[1], face.shape[0]] # x1, y1, x2, y2
|
||||
face = cv2.resize(face, (self.resolution, self.resolution), interpolation=cv2.INTER_LANCZOS4)
|
||||
face = rearrange(torch.from_numpy(face), "h w c -> c h w")
|
||||
return face, box, affine_matrix
|
||||
|
||||
def preprocess_fixed_mask_image(self, image: torch.Tensor, affine_transform=False):
|
||||
if affine_transform:
|
||||
image, _, _ = self.affine_transform(image)
|
||||
else:
|
||||
image = self.resize(image)
|
||||
pixel_values = self.normalize(image / 255.0)
|
||||
masked_pixel_values = pixel_values * self.mask_image
|
||||
return pixel_values, masked_pixel_values, self.mask_image[0:1]
|
||||
|
||||
def prepare_masks_and_masked_images(self, images: Union[torch.Tensor, np.ndarray], affine_transform=False):
|
||||
if isinstance(images, np.ndarray):
|
||||
images = torch.from_numpy(images)
|
||||
if images.shape[3] == 3:
|
||||
images = rearrange(images, "f h w c -> f c h w")
|
||||
if self.mask == "fix_mask":
|
||||
results = [self.preprocess_fixed_mask_image(image, affine_transform=affine_transform) for image in images]
|
||||
else:
|
||||
results = [self.preprocess_one_masked_image(image) for image in images]
|
||||
|
||||
pixel_values_list, masked_pixel_values_list, masks_list = list(zip(*results))
|
||||
return torch.stack(pixel_values_list), torch.stack(masked_pixel_values_list), torch.stack(masks_list)
|
||||
|
||||
def process_images(self, images: Union[torch.Tensor, np.ndarray]):
|
||||
if isinstance(images, np.ndarray):
|
||||
images = torch.from_numpy(images)
|
||||
if images.shape[3] == 3:
|
||||
images = rearrange(images, "f h w c -> f c h w")
|
||||
images = self.resize(images)
|
||||
pixel_values = self.normalize(images / 255.0)
|
||||
return pixel_values
|
||||
|
||||
def close(self):
|
||||
if self.face_mesh is not None:
|
||||
self.face_mesh.close()
|
||||
|
||||
|
||||
def mediapipe_lm478_to_face_alignment_lm68(lm478, return_2d=True):
|
||||
"""
|
||||
lm478: [B, 478, 3] or [478,3]
|
||||
"""
|
||||
# lm478[..., 0] *= W
|
||||
# lm478[..., 1] *= H
|
||||
landmarks_extracted = []
|
||||
for index in landmark_points_68:
|
||||
x = lm478[index][0]
|
||||
y = lm478[index][1]
|
||||
landmarks_extracted.append((x, y))
|
||||
return np.array(landmarks_extracted)
|
||||
|
||||
|
||||
landmark_points_68 = [
|
||||
162,
|
||||
234,
|
||||
93,
|
||||
58,
|
||||
172,
|
||||
136,
|
||||
149,
|
||||
148,
|
||||
152,
|
||||
377,
|
||||
378,
|
||||
365,
|
||||
397,
|
||||
288,
|
||||
323,
|
||||
454,
|
||||
389,
|
||||
71,
|
||||
63,
|
||||
105,
|
||||
66,
|
||||
107,
|
||||
336,
|
||||
296,
|
||||
334,
|
||||
293,
|
||||
301,
|
||||
168,
|
||||
197,
|
||||
5,
|
||||
4,
|
||||
75,
|
||||
97,
|
||||
2,
|
||||
326,
|
||||
305,
|
||||
33,
|
||||
160,
|
||||
158,
|
||||
133,
|
||||
153,
|
||||
144,
|
||||
362,
|
||||
385,
|
||||
387,
|
||||
263,
|
||||
373,
|
||||
380,
|
||||
61,
|
||||
39,
|
||||
37,
|
||||
0,
|
||||
267,
|
||||
269,
|
||||
291,
|
||||
405,
|
||||
314,
|
||||
17,
|
||||
84,
|
||||
181,
|
||||
78,
|
||||
82,
|
||||
13,
|
||||
312,
|
||||
308,
|
||||
317,
|
||||
14,
|
||||
87,
|
||||
]
|
||||
|
||||
|
||||
# Refer to https://storage.googleapis.com/mediapipe-assets/documentation/mediapipe_face_landmark_fullsize.png
|
||||
mouth_surround_landmarks = [
|
||||
164,
|
||||
165,
|
||||
167,
|
||||
92,
|
||||
186,
|
||||
57,
|
||||
43,
|
||||
106,
|
||||
182,
|
||||
83,
|
||||
18,
|
||||
313,
|
||||
406,
|
||||
335,
|
||||
273,
|
||||
287,
|
||||
410,
|
||||
322,
|
||||
391,
|
||||
393,
|
||||
]
|
||||
|
||||
face_surround_landmarks = [
|
||||
152,
|
||||
377,
|
||||
400,
|
||||
378,
|
||||
379,
|
||||
365,
|
||||
397,
|
||||
288,
|
||||
435,
|
||||
433,
|
||||
411,
|
||||
425,
|
||||
423,
|
||||
327,
|
||||
326,
|
||||
94,
|
||||
97,
|
||||
98,
|
||||
203,
|
||||
205,
|
||||
187,
|
||||
213,
|
||||
215,
|
||||
58,
|
||||
172,
|
||||
136,
|
||||
150,
|
||||
149,
|
||||
176,
|
||||
148,
|
||||
]
|
||||
|
||||
if __name__ == "__main__":
|
||||
image_processor = ImageProcessor(512, mask="fix_mask")
|
||||
video = cv2.VideoCapture("assets/demo1_video.mp4")
|
||||
while True:
|
||||
ret, frame = video.read()
|
||||
# if not ret:
|
||||
# break
|
||||
|
||||
# cv2.imwrite("image.jpg", frame)
|
||||
|
||||
frame = rearrange(torch.Tensor(frame).type(torch.uint8), "h w c -> c h w")
|
||||
# face, masked_face, _ = image_processor.preprocess_fixed_mask_image(frame, affine_transform=True)
|
||||
face, _, _ = image_processor.affine_transform(frame)
|
||||
|
||||
break
|
||||
|
||||
face = (rearrange(face, "c h w -> h w c").detach().cpu().numpy()).astype(np.uint8)
|
||||
cv2.imwrite("face.jpg", face)
|
||||
|
||||
# masked_face = (rearrange(masked_face, "c h w -> h w c").detach().cpu().numpy()).astype(np.uint8)
|
||||
# cv2.imwrite("masked_face.jpg", masked_face)
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 1.2 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 1.1 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 1.2 KiB |
+260
-378
@@ -1,378 +1,260 @@
|
||||
# Copyright (c) 2024 Bytedance Ltd. and/or its affiliates
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import os
|
||||
import imageio
|
||||
import numpy as np
|
||||
import json
|
||||
from typing import Union
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torchvision
|
||||
import torch.distributed as dist
|
||||
from torchvision import transforms
|
||||
|
||||
from tqdm import tqdm
|
||||
from einops import rearrange
|
||||
import cv2
|
||||
from decord import AudioReader, VideoReader
|
||||
import shutil
|
||||
import subprocess
|
||||
|
||||
|
||||
# Machine epsilon for a float32 (single precision)
|
||||
eps = np.finfo(np.float32).eps
|
||||
|
||||
|
||||
def read_json(filepath: str):
|
||||
with open(filepath) as f:
|
||||
json_dict = json.load(f)
|
||||
return json_dict
|
||||
|
||||
|
||||
def read_video(video_path: str, change_fps=True, use_decord=True):
|
||||
if change_fps:
|
||||
# Create temp directory next to the video file
|
||||
temp_dir = os.path.join(os.path.dirname(video_path), "temp")
|
||||
if os.path.exists(temp_dir):
|
||||
shutil.rmtree(temp_dir)
|
||||
os.makedirs(temp_dir, exist_ok=True)
|
||||
|
||||
# Normalize paths for Windows
|
||||
video_path = os.path.normpath(video_path).replace('\\', '/')
|
||||
temp_video = os.path.join(temp_dir, 'video.mp4').replace('\\', '/')
|
||||
|
||||
# Use proper path quoting for FFmpeg
|
||||
command = f'ffmpeg -loglevel error -y -nostdin -i "{video_path}" -r 25 -crf 18 "{temp_video}"'
|
||||
|
||||
print(f"Running FFmpeg command: {command}")
|
||||
subprocess.run(command, shell=True, check=True)
|
||||
target_video_path = temp_video
|
||||
else:
|
||||
target_video_path = video_path
|
||||
|
||||
print(f"Reading video from: {target_video_path}")
|
||||
|
||||
if use_decord:
|
||||
return read_video_decord(target_video_path)
|
||||
else:
|
||||
return read_video_cv2(target_video_path)
|
||||
|
||||
|
||||
def read_video_decord(video_path: str):
|
||||
vr = VideoReader(video_path)
|
||||
video_frames = vr[:].asnumpy()
|
||||
vr.seek(0)
|
||||
return video_frames
|
||||
|
||||
|
||||
def read_video_cv2(video_path: str):
|
||||
# Open the video file
|
||||
video_path = os.path.normpath(video_path)
|
||||
print(f"Opening video with CV2: {video_path}")
|
||||
cap = cv2.VideoCapture(video_path)
|
||||
|
||||
# Check if the video was opened successfully
|
||||
if not cap.isOpened():
|
||||
print(f"Error: Could not open video at path: {video_path}")
|
||||
return np.array([])
|
||||
|
||||
frames = []
|
||||
frame_count = 0
|
||||
|
||||
while True:
|
||||
# Read a frame
|
||||
ret, frame = cap.read()
|
||||
|
||||
# If frame is read correctly ret is True
|
||||
if not ret:
|
||||
break
|
||||
|
||||
# Convert BGR to RGB
|
||||
frame_rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
frames.append(frame_rgb)
|
||||
frame_count += 1
|
||||
|
||||
# Release the video capture object
|
||||
cap.release()
|
||||
|
||||
print(f"Successfully read {frame_count} frames from video")
|
||||
return np.array(frames)
|
||||
|
||||
|
||||
def read_audio(audio_path: str, audio_sample_rate: int = 16000):
|
||||
if audio_path is None:
|
||||
raise ValueError("Audio path is required.")
|
||||
ar = AudioReader(audio_path, sample_rate=audio_sample_rate, mono=True)
|
||||
|
||||
# To access the audio samples
|
||||
audio_samples = torch.from_numpy(ar[:].asnumpy())
|
||||
audio_samples = audio_samples.squeeze(0)
|
||||
|
||||
return audio_samples
|
||||
|
||||
|
||||
def write_video(video_output_path: str, video_frames: np.ndarray, fps: int):
|
||||
height, width = video_frames[0].shape[:2]
|
||||
out = cv2.VideoWriter(video_output_path, cv2.VideoWriter_fourcc(*"mp4v"), fps, (width, height))
|
||||
# out = cv2.VideoWriter(video_output_path, cv2.VideoWriter_fourcc(*"vp09"), fps, (width, height))
|
||||
for frame in video_frames:
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_RGB2BGR)
|
||||
out.write(frame)
|
||||
out.release()
|
||||
|
||||
|
||||
def init_dist(backend="nccl", **kwargs):
|
||||
"""Initializes distributed environment."""
|
||||
rank = int(os.environ["RANK"])
|
||||
num_gpus = torch.cuda.device_count()
|
||||
if num_gpus == 0:
|
||||
raise RuntimeError("No GPUs available for training.")
|
||||
local_rank = rank % num_gpus
|
||||
torch.cuda.set_device(local_rank)
|
||||
dist.init_process_group(backend=backend, **kwargs)
|
||||
|
||||
return local_rank
|
||||
|
||||
|
||||
def zero_rank_print(s):
|
||||
if dist.is_initialized() and dist.get_rank() == 0:
|
||||
print("### " + s)
|
||||
|
||||
|
||||
def zero_rank_log(logger, message: str):
|
||||
if dist.is_initialized() and dist.get_rank() == 0:
|
||||
logger.info(message)
|
||||
|
||||
|
||||
def make_audio_window(audio_embeddings: torch.Tensor, window_size: int):
|
||||
audio_window = []
|
||||
end_idx = audio_embeddings.shape[1] - window_size + 1
|
||||
for i in range(end_idx):
|
||||
audio_window.append(audio_embeddings[:, i : i + window_size, :])
|
||||
audio_window = torch.stack(audio_window)
|
||||
audio_window = rearrange(audio_window, "f b w d -> b f w d")
|
||||
return audio_window
|
||||
|
||||
|
||||
def check_video_fps(video_path: str):
|
||||
cam = cv2.VideoCapture(video_path)
|
||||
fps = cam.get(cv2.CAP_PROP_FPS)
|
||||
if fps != 25:
|
||||
raise ValueError(f"Video FPS is not 25, it is {fps}. Please convert the video to 25 FPS.")
|
||||
|
||||
|
||||
def tailor_tensor_to_length(tensor: torch.Tensor, length: int):
|
||||
if len(tensor) == length:
|
||||
return tensor
|
||||
elif len(tensor) > length:
|
||||
return tensor[:length]
|
||||
else:
|
||||
return torch.cat([tensor, tensor[-1].repeat(length - len(tensor))])
|
||||
|
||||
|
||||
def save_videos_grid(videos: torch.Tensor, path: str, rescale=False, n_rows=6, fps=8):
|
||||
videos = rearrange(videos, "b c f h w -> f b c h w")
|
||||
outputs = []
|
||||
for x in videos:
|
||||
x = torchvision.utils.make_grid(x, nrow=n_rows)
|
||||
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
|
||||
if rescale:
|
||||
x = (x + 1.0) / 2.0 # -1,1 -> 0,1
|
||||
x = (x * 255).numpy().astype(np.uint8)
|
||||
outputs.append(x)
|
||||
|
||||
os.makedirs(os.path.dirname(path), exist_ok=True)
|
||||
imageio.mimsave(path, outputs, fps=fps)
|
||||
|
||||
|
||||
def interpolate_features(features: torch.Tensor, output_len: int) -> torch.Tensor:
|
||||
features = features.cpu().numpy()
|
||||
input_len, num_features = features.shape
|
||||
|
||||
input_timesteps = np.linspace(0, 10, input_len)
|
||||
output_timesteps = np.linspace(0, 10, output_len)
|
||||
output_features = np.zeros((output_len, num_features))
|
||||
for feat in range(num_features):
|
||||
output_features[:, feat] = np.interp(output_timesteps, input_timesteps, features[:, feat])
|
||||
return torch.from_numpy(output_features)
|
||||
|
||||
|
||||
# DDIM Inversion
|
||||
@torch.no_grad()
|
||||
def init_prompt(prompt, pipeline):
|
||||
uncond_input = pipeline.tokenizer(
|
||||
[""], padding="max_length", max_length=pipeline.tokenizer.model_max_length, return_tensors="pt"
|
||||
)
|
||||
uncond_embeddings = pipeline.text_encoder(uncond_input.input_ids.to(pipeline.device))[0]
|
||||
text_input = pipeline.tokenizer(
|
||||
[prompt],
|
||||
padding="max_length",
|
||||
max_length=pipeline.tokenizer.model_max_length,
|
||||
truncation=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
text_embeddings = pipeline.text_encoder(text_input.input_ids.to(pipeline.device))[0]
|
||||
context = torch.cat([uncond_embeddings, text_embeddings])
|
||||
|
||||
return context
|
||||
|
||||
|
||||
def reversed_forward(ddim_scheduler, pred_noise, timesteps, x_t):
|
||||
# Compute alphas, betas
|
||||
alpha_prod_t = ddim_scheduler.alphas_cumprod[timesteps]
|
||||
beta_prod_t = 1 - alpha_prod_t
|
||||
|
||||
# 3. compute predicted original sample from predicted noise also called
|
||||
# "predicted x_0" of formula (12) from https://arxiv.org/pdf/2010.02502.pdf
|
||||
if ddim_scheduler.config.prediction_type == "epsilon":
|
||||
beta_prod_t = beta_prod_t[:, None, None, None, None]
|
||||
alpha_prod_t = alpha_prod_t[:, None, None, None, None]
|
||||
pred_original_sample = (x_t - beta_prod_t ** (0.5) * pred_noise) / alpha_prod_t ** (0.5)
|
||||
else:
|
||||
raise NotImplementedError("This prediction type is not implemented yet")
|
||||
|
||||
# Clip "predicted x_0"
|
||||
if ddim_scheduler.config.clip_sample:
|
||||
pred_original_sample = torch.clamp(pred_original_sample, -1, 1)
|
||||
return pred_original_sample
|
||||
|
||||
|
||||
def next_step(
|
||||
model_output: Union[torch.FloatTensor, np.ndarray],
|
||||
timestep: int,
|
||||
sample: Union[torch.FloatTensor, np.ndarray],
|
||||
ddim_scheduler,
|
||||
):
|
||||
timestep, next_timestep = (
|
||||
min(timestep - ddim_scheduler.config.num_train_timesteps // ddim_scheduler.num_inference_steps, 999),
|
||||
timestep,
|
||||
)
|
||||
alpha_prod_t = ddim_scheduler.alphas_cumprod[timestep] if timestep >= 0 else ddim_scheduler.final_alpha_cumprod
|
||||
alpha_prod_t_next = ddim_scheduler.alphas_cumprod[next_timestep]
|
||||
beta_prod_t = 1 - alpha_prod_t
|
||||
next_original_sample = (sample - beta_prod_t**0.5 * model_output) / alpha_prod_t**0.5
|
||||
next_sample_direction = (1 - alpha_prod_t_next) ** 0.5 * model_output
|
||||
next_sample = alpha_prod_t_next**0.5 * next_original_sample + next_sample_direction
|
||||
return next_sample
|
||||
|
||||
|
||||
def get_noise_pred_single(latents, t, context, unet):
|
||||
noise_pred = unet(latents, t, encoder_hidden_states=context)["sample"]
|
||||
return noise_pred
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def ddim_loop(pipeline, ddim_scheduler, latent, num_inv_steps, prompt):
|
||||
context = init_prompt(prompt, pipeline)
|
||||
uncond_embeddings, cond_embeddings = context.chunk(2)
|
||||
all_latent = [latent]
|
||||
latent = latent.clone().detach()
|
||||
for i in tqdm(range(num_inv_steps)):
|
||||
t = ddim_scheduler.timesteps[len(ddim_scheduler.timesteps) - i - 1]
|
||||
noise_pred = get_noise_pred_single(latent, t, cond_embeddings, pipeline.unet)
|
||||
latent = next_step(noise_pred, t, latent, ddim_scheduler)
|
||||
all_latent.append(latent)
|
||||
return all_latent
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def ddim_inversion(pipeline, ddim_scheduler, video_latent, num_inv_steps, prompt=""):
|
||||
ddim_latents = ddim_loop(pipeline, ddim_scheduler, video_latent, num_inv_steps, prompt)
|
||||
return ddim_latents
|
||||
|
||||
|
||||
def plot_loss_chart(save_path: str, *args):
|
||||
# Creating the plot
|
||||
plt.figure()
|
||||
for loss_line in args:
|
||||
plt.plot(loss_line[1], loss_line[2], label=loss_line[0])
|
||||
plt.xlabel("Step")
|
||||
plt.ylabel("Loss")
|
||||
plt.legend()
|
||||
|
||||
# Save the figure to a file
|
||||
plt.savefig(save_path)
|
||||
|
||||
# Close the figure to free memory
|
||||
plt.close()
|
||||
|
||||
|
||||
CRED = "\033[91m"
|
||||
CEND = "\033[0m"
|
||||
|
||||
|
||||
def red_text(text: str):
|
||||
return f"{CRED}{text}{CEND}"
|
||||
|
||||
|
||||
log_loss = nn.BCELoss(reduction="none")
|
||||
|
||||
|
||||
def cosine_loss(vision_embeds, audio_embeds, y):
|
||||
sims = nn.functional.cosine_similarity(vision_embeds, audio_embeds)
|
||||
# sims[sims!=sims] = 0 # remove nan
|
||||
# sims = sims.clamp(0, 1)
|
||||
loss = log_loss(sims.unsqueeze(1), y).squeeze()
|
||||
return loss
|
||||
|
||||
|
||||
def save_image(image, save_path):
|
||||
# input size (C, H, W)
|
||||
image = (image / 2 + 0.5).clamp(0, 1)
|
||||
image = (image * 255).to(torch.uint8)
|
||||
image = transforms.ToPILImage()(image)
|
||||
# Save the image copy
|
||||
image.save(save_path)
|
||||
|
||||
# Close the image file
|
||||
image.close()
|
||||
|
||||
|
||||
def gather_loss(loss, device):
|
||||
# Sum the local loss across all processes
|
||||
local_loss = loss.item()
|
||||
global_loss = torch.tensor(local_loss, dtype=torch.float32).to(device)
|
||||
dist.all_reduce(global_loss, op=dist.ReduceOp.SUM)
|
||||
|
||||
# Calculate the average loss across all processes
|
||||
global_average_loss = global_loss.item() / dist.get_world_size()
|
||||
return global_average_loss
|
||||
|
||||
|
||||
def gather_video_paths_recursively(input_dir):
|
||||
print(f"Recursively gathering video paths of {input_dir} ...")
|
||||
paths = []
|
||||
gather_video_paths(input_dir, paths)
|
||||
return paths
|
||||
|
||||
|
||||
def gather_video_paths(input_dir, paths):
|
||||
for file in sorted(os.listdir(input_dir)):
|
||||
if file.endswith(".mp4"):
|
||||
filepath = os.path.join(input_dir, file)
|
||||
paths.append(filepath)
|
||||
elif os.path.isdir(os.path.join(input_dir, file)):
|
||||
gather_video_paths(os.path.join(input_dir, file), paths)
|
||||
|
||||
|
||||
def count_video_time(video_path):
|
||||
video = cv2.VideoCapture(video_path)
|
||||
|
||||
frame_count = video.get(cv2.CAP_PROP_FRAME_COUNT)
|
||||
fps = video.get(cv2.CAP_PROP_FPS)
|
||||
return frame_count / fps
|
||||
# Copyright (c) 2024 Bytedance Ltd. and/or its affiliates
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import os
|
||||
import imageio
|
||||
import numpy as np
|
||||
import json
|
||||
from typing import Union
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torchvision
|
||||
import torch.distributed as dist
|
||||
from torchvision import transforms
|
||||
|
||||
from einops import rearrange
|
||||
import cv2
|
||||
from decord import AudioReader, VideoReader
|
||||
import shutil
|
||||
import subprocess
|
||||
|
||||
|
||||
# Machine epsilon for a float32 (single precision)
|
||||
eps = np.finfo(np.float32).eps
|
||||
|
||||
|
||||
def read_json(filepath: str):
|
||||
with open(filepath) as f:
|
||||
json_dict = json.load(f)
|
||||
return json_dict
|
||||
|
||||
|
||||
def read_video(video_path: str, change_fps=True, use_decord=True):
|
||||
if change_fps:
|
||||
temp_dir = "temp"
|
||||
if os.path.exists(temp_dir):
|
||||
shutil.rmtree(temp_dir)
|
||||
os.makedirs(temp_dir, exist_ok=True)
|
||||
command = (
|
||||
f"ffmpeg -loglevel error -y -nostdin -i {video_path} -r 25 -crf 18 {os.path.join(temp_dir, 'video.mp4')}"
|
||||
)
|
||||
subprocess.run(command, shell=True)
|
||||
target_video_path = os.path.join(temp_dir, "video.mp4")
|
||||
else:
|
||||
target_video_path = video_path
|
||||
|
||||
if use_decord:
|
||||
return read_video_decord(target_video_path)
|
||||
else:
|
||||
return read_video_cv2(target_video_path)
|
||||
|
||||
|
||||
def read_video_decord(video_path: str):
|
||||
vr = VideoReader(video_path)
|
||||
video_frames = vr[:].asnumpy()
|
||||
vr.seek(0)
|
||||
return video_frames
|
||||
|
||||
|
||||
def read_video_cv2(video_path: str):
|
||||
# Open the video file
|
||||
cap = cv2.VideoCapture(video_path)
|
||||
|
||||
# Check if the video was opened successfully
|
||||
if not cap.isOpened():
|
||||
print("Error: Could not open video.")
|
||||
return np.array([])
|
||||
|
||||
frames = []
|
||||
|
||||
while True:
|
||||
# Read a frame
|
||||
ret, frame = cap.read()
|
||||
|
||||
# If frame is read correctly ret is True
|
||||
if not ret:
|
||||
break
|
||||
|
||||
# Convert BGR to RGB
|
||||
frame_rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
|
||||
frames.append(frame_rgb)
|
||||
|
||||
# Release the video capture object
|
||||
cap.release()
|
||||
|
||||
return np.array(frames)
|
||||
|
||||
|
||||
def read_audio(audio_path: str, audio_sample_rate: int = 16000):
|
||||
if audio_path is None:
|
||||
raise ValueError("Audio path is required.")
|
||||
ar = AudioReader(audio_path, sample_rate=audio_sample_rate, mono=True)
|
||||
|
||||
# To access the audio samples
|
||||
audio_samples = torch.from_numpy(ar[:].asnumpy())
|
||||
audio_samples = audio_samples.squeeze(0)
|
||||
|
||||
return audio_samples
|
||||
|
||||
|
||||
def write_video(video_output_path: str, video_frames: np.ndarray, fps: int):
|
||||
height, width = video_frames[0].shape[:2]
|
||||
out = cv2.VideoWriter(video_output_path, cv2.VideoWriter_fourcc(*"mp4v"), fps, (width, height))
|
||||
# out = cv2.VideoWriter(video_output_path, cv2.VideoWriter_fourcc(*"vp09"), fps, (width, height))
|
||||
for frame in video_frames:
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_RGB2BGR)
|
||||
out.write(frame)
|
||||
out.release()
|
||||
|
||||
|
||||
def init_dist(backend="nccl", **kwargs):
|
||||
"""Initializes distributed environment."""
|
||||
rank = int(os.environ["RANK"])
|
||||
num_gpus = torch.cuda.device_count()
|
||||
if num_gpus == 0:
|
||||
raise RuntimeError("No GPUs available for training.")
|
||||
local_rank = rank % num_gpus
|
||||
torch.cuda.set_device(local_rank)
|
||||
dist.init_process_group(backend=backend, **kwargs)
|
||||
|
||||
return local_rank
|
||||
|
||||
|
||||
def zero_rank_print(s):
|
||||
if dist.is_initialized() and dist.get_rank() == 0:
|
||||
print("### " + s)
|
||||
|
||||
|
||||
def zero_rank_log(logger, message: str):
|
||||
if dist.is_initialized() and dist.get_rank() == 0:
|
||||
logger.info(message)
|
||||
|
||||
|
||||
def check_video_fps(video_path: str):
|
||||
cam = cv2.VideoCapture(video_path)
|
||||
fps = cam.get(cv2.CAP_PROP_FPS)
|
||||
if fps != 25:
|
||||
raise ValueError(f"Video FPS is not 25, it is {fps}. Please convert the video to 25 FPS.")
|
||||
|
||||
|
||||
def one_step_sampling(ddim_scheduler, pred_noise, timesteps, x_t):
|
||||
# Compute alphas, betas
|
||||
alpha_prod_t = ddim_scheduler.alphas_cumprod[timesteps].to(dtype=pred_noise.dtype)
|
||||
beta_prod_t = 1 - alpha_prod_t
|
||||
|
||||
# 3. compute predicted original sample from predicted noise also called
|
||||
# "predicted x_0" of formula (12) from https://arxiv.org/abs/2010.02502
|
||||
if ddim_scheduler.config.prediction_type == "epsilon":
|
||||
beta_prod_t = beta_prod_t[:, None, None, None, None]
|
||||
alpha_prod_t = alpha_prod_t[:, None, None, None, None]
|
||||
pred_original_sample = (x_t - beta_prod_t ** (0.5) * pred_noise) / alpha_prod_t ** (0.5)
|
||||
else:
|
||||
raise NotImplementedError("This prediction type is not implemented yet")
|
||||
|
||||
# Clip "predicted x_0"
|
||||
if ddim_scheduler.config.clip_sample:
|
||||
pred_original_sample = torch.clamp(pred_original_sample, -1, 1)
|
||||
return pred_original_sample
|
||||
|
||||
|
||||
def plot_loss_chart(save_path: str, *args):
|
||||
# Creating the plot
|
||||
plt.figure()
|
||||
for loss_line in args:
|
||||
plt.plot(loss_line[1], loss_line[2], label=loss_line[0])
|
||||
plt.xlabel("Step")
|
||||
plt.ylabel("Loss")
|
||||
plt.legend()
|
||||
|
||||
# Save the figure to a file
|
||||
plt.savefig(save_path)
|
||||
|
||||
# Close the figure to free memory
|
||||
plt.close()
|
||||
|
||||
|
||||
CRED = "\033[91m"
|
||||
CEND = "\033[0m"
|
||||
|
||||
|
||||
def red_text(text: str):
|
||||
return f"{CRED}{text}{CEND}"
|
||||
|
||||
|
||||
log_loss = nn.BCELoss(reduction="none")
|
||||
|
||||
|
||||
def cosine_loss(vision_embeds, audio_embeds, y):
|
||||
sims = nn.functional.cosine_similarity(vision_embeds, audio_embeds)
|
||||
# sims[sims!=sims] = 0 # remove nan
|
||||
# sims = sims.clamp(0, 1)
|
||||
loss = log_loss(sims.unsqueeze(1), y).squeeze()
|
||||
return loss
|
||||
|
||||
|
||||
def save_image(image, save_path):
|
||||
# input size (C, H, W)
|
||||
image = (image / 2 + 0.5).clamp(0, 1)
|
||||
image = (image * 255).to(torch.uint8)
|
||||
image = transforms.ToPILImage()(image)
|
||||
# Save the image copy
|
||||
image.save(save_path)
|
||||
|
||||
# Close the image file
|
||||
image.close()
|
||||
|
||||
|
||||
def gather_loss(loss, device):
|
||||
# Sum the local loss across all processes
|
||||
local_loss = loss.item()
|
||||
global_loss = torch.tensor(local_loss, dtype=torch.float32).to(device)
|
||||
dist.all_reduce(global_loss, op=dist.ReduceOp.SUM)
|
||||
|
||||
# Calculate the average loss across all processes
|
||||
global_average_loss = global_loss.item() / dist.get_world_size()
|
||||
return global_average_loss
|
||||
|
||||
|
||||
def gather_video_paths_recursively(input_dir):
|
||||
print(f"Recursively gathering video paths of {input_dir} ...")
|
||||
paths = []
|
||||
gather_video_paths(input_dir, paths)
|
||||
return paths
|
||||
|
||||
|
||||
def gather_video_paths(input_dir, paths):
|
||||
for file in sorted(os.listdir(input_dir)):
|
||||
if file.endswith(".mp4"):
|
||||
filepath = os.path.join(input_dir, file)
|
||||
paths.append(filepath)
|
||||
elif os.path.isdir(os.path.join(input_dir, file)):
|
||||
gather_video_paths(os.path.join(input_dir, file), paths)
|
||||
|
||||
|
||||
def count_video_time(video_path):
|
||||
video = cv2.VideoCapture(video_path)
|
||||
|
||||
frame_count = video.get(cv2.CAP_PROP_FRAME_COUNT)
|
||||
fps = video.get(cv2.CAP_PROP_FPS)
|
||||
return frame_count / fps
|
||||
|
||||
|
||||
def check_ffmpeg_installed():
|
||||
# Run the ffmpeg command with the -version argument to check if it's installed
|
||||
result = subprocess.run("ffmpeg -version", stdout=subprocess.PIPE, stderr=subprocess.PIPE, shell=True)
|
||||
if not result.returncode == 0:
|
||||
raise FileNotFoundError("ffmpeg not found, please install it by:\n $ conda install -c conda-forge ffmpeg")
|
||||
|
||||
+162
-166
@@ -1,166 +1,162 @@
|
||||
# Adapted from https://github.com/TMElyralab/MuseTalk/blob/main/musetalk/whisper/audio2feature.py
|
||||
|
||||
from .whisper import load_model
|
||||
import numpy as np
|
||||
import torch
|
||||
import os
|
||||
|
||||
|
||||
class Audio2Feature:
|
||||
def __init__(
|
||||
self,
|
||||
model_path="checkpoints/whisper/tiny.pt",
|
||||
device=None,
|
||||
audio_cache_dir=None,
|
||||
num_frames=16,
|
||||
):
|
||||
self.model = load_model(model_path, device)
|
||||
self.audio_cache_dir = audio_cache_dir
|
||||
self.num_frames = num_frames
|
||||
self.embedding_dim = self.model.dims.n_audio_state
|
||||
|
||||
def get_sliced_feature(self, feature_array, vid_idx, audio_feat_length=[2, 2], fps=25):
|
||||
"""
|
||||
Get sliced features based on a given index
|
||||
:param feature_array:
|
||||
:param start_idx: the start index of the feature
|
||||
:param audio_feat_length:
|
||||
:return:
|
||||
"""
|
||||
length = len(feature_array)
|
||||
selected_feature = []
|
||||
selected_idx = []
|
||||
|
||||
center_idx = int(vid_idx * 50 / fps)
|
||||
left_idx = center_idx - audio_feat_length[0] * 2
|
||||
right_idx = center_idx + (audio_feat_length[1] + 1) * 2
|
||||
|
||||
for idx in range(left_idx, right_idx):
|
||||
idx = max(0, idx)
|
||||
idx = min(length - 1, idx)
|
||||
x = feature_array[idx]
|
||||
selected_feature.append(x)
|
||||
selected_idx.append(idx)
|
||||
|
||||
selected_feature = torch.cat(selected_feature, dim=0)
|
||||
selected_feature = selected_feature.reshape(-1, self.embedding_dim) # 50*384
|
||||
return selected_feature, selected_idx
|
||||
|
||||
def get_sliced_feature_sparse(self, feature_array, vid_idx, audio_feat_length=[2, 2], fps=25):
|
||||
"""
|
||||
Get sliced features based on a given index
|
||||
:param feature_array:
|
||||
:param start_idx: the start index of the feature
|
||||
:param audio_feat_length:
|
||||
:return:
|
||||
"""
|
||||
length = len(feature_array)
|
||||
selected_feature = []
|
||||
selected_idx = []
|
||||
|
||||
for dt in range(-audio_feat_length[0], audio_feat_length[1] + 1):
|
||||
left_idx = int((vid_idx + dt) * 50 / fps)
|
||||
if left_idx < 1 or left_idx > length - 1:
|
||||
left_idx = max(0, left_idx)
|
||||
left_idx = min(length - 1, left_idx)
|
||||
|
||||
x = feature_array[left_idx]
|
||||
x = x[np.newaxis, :, :]
|
||||
x = np.repeat(x, 2, axis=0)
|
||||
selected_feature.append(x)
|
||||
selected_idx.append(left_idx)
|
||||
selected_idx.append(left_idx)
|
||||
else:
|
||||
x = feature_array[left_idx - 1 : left_idx + 1]
|
||||
selected_feature.append(x)
|
||||
selected_idx.append(left_idx - 1)
|
||||
selected_idx.append(left_idx)
|
||||
selected_feature = np.concatenate(selected_feature, axis=0)
|
||||
selected_feature = selected_feature.reshape(-1, self.embedding_dim) # 50*384
|
||||
selected_feature = torch.from_numpy(selected_feature)
|
||||
return selected_feature, selected_idx
|
||||
|
||||
def feature2chunks(self, feature_array, fps, audio_feat_length=[2, 2]):
|
||||
whisper_chunks = []
|
||||
whisper_idx_multiplier = 50.0 / fps
|
||||
i = 0
|
||||
print(f"video in {fps} FPS, audio idx in 50FPS")
|
||||
|
||||
while True:
|
||||
start_idx = int(i * whisper_idx_multiplier)
|
||||
selected_feature, selected_idx = self.get_sliced_feature(
|
||||
feature_array=feature_array, vid_idx=i, audio_feat_length=audio_feat_length, fps=fps
|
||||
)
|
||||
# print(f"i:{i},selected_idx {selected_idx}")
|
||||
whisper_chunks.append(selected_feature)
|
||||
i += 1
|
||||
if start_idx > len(feature_array):
|
||||
break
|
||||
|
||||
return whisper_chunks
|
||||
|
||||
def _audio2feat(self, audio_path: str):
|
||||
# get the sample rate of the audio
|
||||
result = self.model.transcribe(audio_path)
|
||||
embed_list = []
|
||||
for emb in result["segments"]:
|
||||
encoder_embeddings = emb["encoder_embeddings"]
|
||||
encoder_embeddings = encoder_embeddings.transpose(0, 2, 1, 3)
|
||||
encoder_embeddings = encoder_embeddings.squeeze(0)
|
||||
start_idx = int(emb["start"])
|
||||
end_idx = int(emb["end"])
|
||||
emb_end_idx = int((end_idx - start_idx) / 2)
|
||||
embed_list.append(encoder_embeddings[:emb_end_idx])
|
||||
concatenated_array = torch.from_numpy(np.concatenate(embed_list, axis=0))
|
||||
return concatenated_array
|
||||
|
||||
def audio2feat(self, audio_path):
|
||||
if self.audio_cache_dir == "" or self.audio_cache_dir is None:
|
||||
return self._audio2feat(audio_path)
|
||||
|
||||
audio_cache_path = os.path.join(self.audio_cache_dir, os.path.basename(audio_path) + ".pt")
|
||||
|
||||
if os.path.isfile(audio_cache_path):
|
||||
try:
|
||||
audio_feat = torch.load(audio_cache_path)
|
||||
except Exception as e:
|
||||
print(f"{type(e).__name__} - {e} - {audio_cache_path}")
|
||||
os.remove(audio_cache_path)
|
||||
audio_feat = self._audio2feat(audio_path)
|
||||
torch.save(audio_feat, audio_cache_path)
|
||||
else:
|
||||
audio_feat = self._audio2feat(audio_path)
|
||||
torch.save(audio_feat, audio_cache_path)
|
||||
|
||||
return audio_feat
|
||||
|
||||
def crop_overlap_audio_window(self, audio_feat, start_index):
|
||||
selected_feature_list = []
|
||||
for i in range(start_index, start_index + self.num_frames):
|
||||
selected_feature, selected_idx = self.get_sliced_feature(
|
||||
feature_array=audio_feat, vid_idx=i, audio_feat_length=[2, 2], fps=25
|
||||
)
|
||||
selected_feature_list.append(selected_feature)
|
||||
mel_overlap = torch.stack(selected_feature_list)
|
||||
return mel_overlap
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
audio_encoder = Audio2Feature(model_path="checkpoints/whisper/tiny.pt")
|
||||
audio_path = "./validation/hdtf_val.mp4"
|
||||
array = audio_encoder.audio2feat(audio_path)
|
||||
print(array.shape)
|
||||
fps = 25
|
||||
whisper_idx_multiplier = 50.0 / fps
|
||||
|
||||
i = 0
|
||||
print(f"video in {fps} FPS, audio idx in 50FPS")
|
||||
while True:
|
||||
start_idx = int(i * whisper_idx_multiplier)
|
||||
selected_feature, selected_idx = audio_encoder.get_sliced_feature(
|
||||
feature_array=array, vid_idx=i, audio_feat_length=[2, 2], fps=fps
|
||||
)
|
||||
print(f"video idx {i},\t audio idx {selected_idx},\t shape {selected_feature.shape}")
|
||||
i += 1
|
||||
if start_idx > len(array):
|
||||
break
|
||||
# Adapted from https://github.com/TMElyralab/MuseTalk/blob/main/musetalk/whisper/audio2feature.py
|
||||
|
||||
from .whisper import load_model
|
||||
import numpy as np
|
||||
import torch
|
||||
import os
|
||||
|
||||
|
||||
class Audio2Feature:
|
||||
def __init__(
|
||||
self,
|
||||
model_path="checkpoints/whisper/tiny.pt",
|
||||
device=None,
|
||||
audio_embeds_cache_dir=None,
|
||||
num_frames=16,
|
||||
audio_feat_length=[2, 2],
|
||||
):
|
||||
self.model = load_model(model_path, device)
|
||||
self.audio_embeds_cache_dir = audio_embeds_cache_dir
|
||||
self.num_frames = num_frames
|
||||
self.embedding_dim = self.model.dims.n_audio_state
|
||||
self.audio_feat_length = audio_feat_length
|
||||
|
||||
def get_sliced_feature(self, feature_array, vid_idx, fps=25):
|
||||
"""
|
||||
Get sliced features based on a given index
|
||||
:param feature_array:
|
||||
:param start_idx: the start index of the feature
|
||||
:param audio_feat_length:
|
||||
:return:
|
||||
"""
|
||||
length = len(feature_array)
|
||||
selected_feature = []
|
||||
selected_idx = []
|
||||
|
||||
center_idx = int(vid_idx * 50 / fps)
|
||||
left_idx = center_idx - self.audio_feat_length[0] * 2
|
||||
right_idx = center_idx + (self.audio_feat_length[1] + 1) * 2
|
||||
|
||||
for idx in range(left_idx, right_idx):
|
||||
idx = max(0, idx)
|
||||
idx = min(length - 1, idx)
|
||||
x = feature_array[idx]
|
||||
selected_feature.append(x)
|
||||
selected_idx.append(idx)
|
||||
|
||||
selected_feature = torch.cat(selected_feature, dim=0)
|
||||
selected_feature = selected_feature.reshape(-1, self.embedding_dim) # 50*384
|
||||
return selected_feature, selected_idx
|
||||
|
||||
def get_sliced_feature_sparse(self, feature_array, vid_idx, fps=25):
|
||||
"""
|
||||
Get sliced features based on a given index
|
||||
:param feature_array:
|
||||
:param start_idx: the start index of the feature
|
||||
:param audio_feat_length:
|
||||
:return:
|
||||
"""
|
||||
length = len(feature_array)
|
||||
selected_feature = []
|
||||
selected_idx = []
|
||||
|
||||
for dt in range(-self.audio_feat_length[0], self.audio_feat_length[1] + 1):
|
||||
left_idx = int((vid_idx + dt) * 50 / fps)
|
||||
if left_idx < 1 or left_idx > length - 1:
|
||||
left_idx = max(0, left_idx)
|
||||
left_idx = min(length - 1, left_idx)
|
||||
|
||||
x = feature_array[left_idx]
|
||||
x = x[np.newaxis, :, :]
|
||||
x = np.repeat(x, 2, axis=0)
|
||||
selected_feature.append(x)
|
||||
selected_idx.append(left_idx)
|
||||
selected_idx.append(left_idx)
|
||||
else:
|
||||
x = feature_array[left_idx - 1 : left_idx + 1]
|
||||
selected_feature.append(x)
|
||||
selected_idx.append(left_idx - 1)
|
||||
selected_idx.append(left_idx)
|
||||
selected_feature = np.concatenate(selected_feature, axis=0)
|
||||
selected_feature = selected_feature.reshape(-1, self.embedding_dim) # 50*384
|
||||
selected_feature = torch.from_numpy(selected_feature)
|
||||
return selected_feature, selected_idx
|
||||
|
||||
def feature2chunks(self, feature_array, fps):
|
||||
whisper_chunks = []
|
||||
whisper_idx_multiplier = 50.0 / fps
|
||||
i = 0
|
||||
print(f"video in {fps} FPS, audio idx in 50FPS")
|
||||
|
||||
while True:
|
||||
start_idx = int(i * whisper_idx_multiplier)
|
||||
selected_feature, selected_idx = self.get_sliced_feature(feature_array=feature_array, vid_idx=i, fps=fps)
|
||||
# print(f"i:{i},selected_idx {selected_idx}")
|
||||
whisper_chunks.append(selected_feature)
|
||||
i += 1
|
||||
if start_idx > len(feature_array):
|
||||
break
|
||||
|
||||
return whisper_chunks
|
||||
|
||||
def _audio2feat(self, audio_path: str):
|
||||
# get the sample rate of the audio
|
||||
result = self.model.transcribe(audio_path)
|
||||
embed_list = []
|
||||
for emb in result["segments"]:
|
||||
encoder_embeddings = emb["encoder_embeddings"]
|
||||
encoder_embeddings = encoder_embeddings.transpose(0, 2, 1, 3)
|
||||
encoder_embeddings = encoder_embeddings.squeeze(0)
|
||||
start_idx = int(emb["start"])
|
||||
end_idx = int(emb["end"])
|
||||
emb_end_idx = int((end_idx - start_idx) / 2)
|
||||
embed_list.append(encoder_embeddings[:emb_end_idx])
|
||||
concatenated_array = torch.from_numpy(np.concatenate(embed_list, axis=0))
|
||||
return concatenated_array
|
||||
|
||||
def audio2feat(self, audio_path):
|
||||
if self.audio_embeds_cache_dir == "" or self.audio_embeds_cache_dir is None:
|
||||
return self._audio2feat(audio_path)
|
||||
|
||||
audio_embeds_cache_path = os.path.join(self.audio_embeds_cache_dir, os.path.basename(audio_path) + ".pt")
|
||||
|
||||
if os.path.isfile(audio_embeds_cache_path):
|
||||
try:
|
||||
audio_feat = torch.load(audio_embeds_cache_path, weights_only=True)
|
||||
except Exception as e:
|
||||
print(f"{type(e).__name__} - {e} - {audio_embeds_cache_path}")
|
||||
os.remove(audio_embeds_cache_path)
|
||||
audio_feat = self._audio2feat(audio_path)
|
||||
torch.save(audio_feat, audio_embeds_cache_path)
|
||||
else:
|
||||
audio_feat = self._audio2feat(audio_path)
|
||||
torch.save(audio_feat, audio_embeds_cache_path)
|
||||
|
||||
return audio_feat
|
||||
|
||||
def crop_overlap_audio_window(self, audio_feat, start_index):
|
||||
selected_feature_list = []
|
||||
for i in range(start_index, start_index + self.num_frames):
|
||||
selected_feature, selected_idx = self.get_sliced_feature(feature_array=audio_feat, vid_idx=i, fps=25)
|
||||
selected_feature_list.append(selected_feature)
|
||||
mel_overlap = torch.stack(selected_feature_list)
|
||||
return mel_overlap
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
audio_encoder = Audio2Feature(model_path="checkpoints/whisper/tiny.pt")
|
||||
audio_path = "assets/demo1_audio.wav"
|
||||
array = audio_encoder.audio2feat(audio_path)
|
||||
print(array.shape)
|
||||
fps = 25
|
||||
whisper_idx_multiplier = 50.0 / fps
|
||||
|
||||
i = 0
|
||||
print(f"video in {fps} FPS, audio idx in 50FPS")
|
||||
while True:
|
||||
start_idx = int(i * whisper_idx_multiplier)
|
||||
selected_feature, selected_idx = audio_encoder.get_sliced_feature(feature_array=array, vid_idx=i, fps=fps)
|
||||
print(f"video idx {i},\t audio idx {selected_idx},\t shape {selected_feature.shape}")
|
||||
i += 1
|
||||
if start_idx > len(array):
|
||||
break
|
||||
|
||||
@@ -1,119 +1,122 @@
|
||||
import hashlib
|
||||
import io
|
||||
import os
|
||||
import urllib
|
||||
import warnings
|
||||
from typing import List, Optional, Union
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from .audio import load_audio, log_mel_spectrogram, pad_or_trim
|
||||
from .decoding import DecodingOptions, DecodingResult, decode, detect_language
|
||||
from .model import Whisper, ModelDimensions
|
||||
from .transcribe import transcribe
|
||||
|
||||
|
||||
_MODELS = {
|
||||
"tiny.en": "https://openaipublic.azureedge.net/main/whisper/models/d3dd57d32accea0b295c96e26691aa14d8822fac7d9d27d5dc00b4ca2826dd03/tiny.en.pt",
|
||||
"tiny": "https://openaipublic.azureedge.net/main/whisper/models/65147644a518d12f04e32d6f3b26facc3f8dd46e5390956a9424a650c0ce22b9/tiny.pt",
|
||||
"base.en": "https://openaipublic.azureedge.net/main/whisper/models/25a8566e1d0c1e2231d1c762132cd20e0f96a85d16145c3a00adf5d1ac670ead/base.en.pt",
|
||||
"base": "https://openaipublic.azureedge.net/main/whisper/models/ed3a0b6b1c0edf879ad9b11b1af5a0e6ab5db9205f891f668f8b0e6c6326e34e/base.pt",
|
||||
"small.en": "https://openaipublic.azureedge.net/main/whisper/models/f953ad0fd29cacd07d5a9eda5624af0f6bcf2258be67c92b79389873d91e0872/small.en.pt",
|
||||
"small": "https://openaipublic.azureedge.net/main/whisper/models/9ecf779972d90ba49c06d968637d720dd632c55bbf19d441fb42bf17a411e794/small.pt",
|
||||
"medium.en": "https://openaipublic.azureedge.net/main/whisper/models/d7440d1dc186f76616474e0ff0b3b6b879abc9d1a4926b7adfa41db2d497ab4f/medium.en.pt",
|
||||
"medium": "https://openaipublic.azureedge.net/main/whisper/models/345ae4da62f9b3d59415adc60127b97c714f32e89e936602e85993674d08dcb1/medium.pt",
|
||||
"large": "https://openaipublic.azureedge.net/main/whisper/models/e4b87e7e0bf463eb8e6956e646f1e277e901512310def2c24bf0e11bd3c28e9a/large.pt",
|
||||
"large-v1": "https://openaipublic.azureedge.net/main/whisper/models/e4b87e7e0bf463eb8e6956e646f1e277e901512310def2c24bf0e11bd3c28e9a/large-v1.pt",
|
||||
"large-v2": "https://openaipublic.azureedge.net/main/whisper/models/81f7c96c852ee8fc832187b0132e569d6c3065a3252ed18e56effd0b6a73e524/large-v2.pt",
|
||||
"large-v3": "https://openaipublic.azureedge.net/main/whisper/models/e5b1a55b89c1367dacf97e3e19bfd829a01529dbfdeefa8caeb59b3f1b81dadb/large-v3.pt",
|
||||
}
|
||||
|
||||
|
||||
def _download(url: str, root: str, in_memory: bool) -> Union[bytes, str]:
|
||||
os.makedirs(root, exist_ok=True)
|
||||
|
||||
expected_sha256 = url.split("/")[-2]
|
||||
download_target = os.path.join(root, os.path.basename(url))
|
||||
|
||||
if os.path.exists(download_target) and not os.path.isfile(download_target):
|
||||
raise RuntimeError(f"{download_target} exists and is not a regular file")
|
||||
|
||||
if os.path.isfile(download_target):
|
||||
model_bytes = open(download_target, "rb").read()
|
||||
if hashlib.sha256(model_bytes).hexdigest() == expected_sha256:
|
||||
return model_bytes if in_memory else download_target
|
||||
else:
|
||||
warnings.warn(f"{download_target} exists, but the SHA256 checksum does not match; re-downloading the file")
|
||||
|
||||
with urllib.request.urlopen(url) as source, open(download_target, "wb") as output:
|
||||
with tqdm(
|
||||
total=int(source.info().get("Content-Length")), ncols=80, unit="iB", unit_scale=True, unit_divisor=1024
|
||||
) as loop:
|
||||
while True:
|
||||
buffer = source.read(8192)
|
||||
if not buffer:
|
||||
break
|
||||
|
||||
output.write(buffer)
|
||||
loop.update(len(buffer))
|
||||
|
||||
model_bytes = open(download_target, "rb").read()
|
||||
if hashlib.sha256(model_bytes).hexdigest() != expected_sha256:
|
||||
raise RuntimeError(
|
||||
"Model has been downloaded but the SHA256 checksum does not not match. Please retry loading the model."
|
||||
)
|
||||
|
||||
return model_bytes if in_memory else download_target
|
||||
|
||||
|
||||
def available_models() -> List[str]:
|
||||
"""Returns the names of available models"""
|
||||
return list(_MODELS.keys())
|
||||
|
||||
|
||||
def load_model(
|
||||
name: str, device: Optional[Union[str, torch.device]] = None, download_root: str = None, in_memory: bool = False
|
||||
) -> Whisper:
|
||||
"""
|
||||
Load a Whisper ASR model
|
||||
|
||||
Parameters
|
||||
----------
|
||||
name : str
|
||||
one of the official model names listed by `whisper.available_models()`, or
|
||||
path to a model checkpoint containing the model dimensions and the model state_dict.
|
||||
device : Union[str, torch.device]
|
||||
the PyTorch device to put the model into
|
||||
download_root: str
|
||||
path to download the model files; by default, it uses "~/.cache/whisper"
|
||||
in_memory: bool
|
||||
whether to preload the model weights into host memory
|
||||
|
||||
Returns
|
||||
-------
|
||||
model : Whisper
|
||||
The Whisper ASR model instance
|
||||
"""
|
||||
|
||||
if device is None:
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
if download_root is None:
|
||||
download_root = os.getenv("XDG_CACHE_HOME", os.path.join(os.path.expanduser("~"), ".cache", "whisper"))
|
||||
|
||||
if name in _MODELS:
|
||||
checkpoint_file = _download(_MODELS[name], download_root, in_memory)
|
||||
elif os.path.isfile(name):
|
||||
checkpoint_file = open(name, "rb").read() if in_memory else name
|
||||
else:
|
||||
raise RuntimeError(f"Model {name} not found; available models = {available_models()}")
|
||||
|
||||
with (io.BytesIO(checkpoint_file) if in_memory else open(checkpoint_file, "rb")) as fp:
|
||||
checkpoint = torch.load(fp, map_location=device)
|
||||
del checkpoint_file
|
||||
|
||||
dims = ModelDimensions(**checkpoint["dims"])
|
||||
model = Whisper(dims)
|
||||
model.load_state_dict(checkpoint["model_state_dict"])
|
||||
|
||||
return model.to(device)
|
||||
import hashlib
|
||||
import io
|
||||
import os
|
||||
import urllib
|
||||
import warnings
|
||||
from typing import List, Optional, Union
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from .audio import load_audio, log_mel_spectrogram, pad_or_trim
|
||||
from .decoding import DecodingOptions, DecodingResult, decode, detect_language
|
||||
from .model import Whisper, ModelDimensions
|
||||
from .transcribe import transcribe
|
||||
|
||||
|
||||
_MODELS = {
|
||||
"tiny.en": "https://openaipublic.azureedge.net/main/whisper/models/d3dd57d32accea0b295c96e26691aa14d8822fac7d9d27d5dc00b4ca2826dd03/tiny.en.pt",
|
||||
"tiny": "https://openaipublic.azureedge.net/main/whisper/models/65147644a518d12f04e32d6f3b26facc3f8dd46e5390956a9424a650c0ce22b9/tiny.pt",
|
||||
"base.en": "https://openaipublic.azureedge.net/main/whisper/models/25a8566e1d0c1e2231d1c762132cd20e0f96a85d16145c3a00adf5d1ac670ead/base.en.pt",
|
||||
"base": "https://openaipublic.azureedge.net/main/whisper/models/ed3a0b6b1c0edf879ad9b11b1af5a0e6ab5db9205f891f668f8b0e6c6326e34e/base.pt",
|
||||
"small.en": "https://openaipublic.azureedge.net/main/whisper/models/f953ad0fd29cacd07d5a9eda5624af0f6bcf2258be67c92b79389873d91e0872/small.en.pt",
|
||||
"small": "https://openaipublic.azureedge.net/main/whisper/models/9ecf779972d90ba49c06d968637d720dd632c55bbf19d441fb42bf17a411e794/small.pt",
|
||||
"medium.en": "https://openaipublic.azureedge.net/main/whisper/models/d7440d1dc186f76616474e0ff0b3b6b879abc9d1a4926b7adfa41db2d497ab4f/medium.en.pt",
|
||||
"medium": "https://openaipublic.azureedge.net/main/whisper/models/345ae4da62f9b3d59415adc60127b97c714f32e89e936602e85993674d08dcb1/medium.pt",
|
||||
"large": "https://openaipublic.azureedge.net/main/whisper/models/e4b87e7e0bf463eb8e6956e646f1e277e901512310def2c24bf0e11bd3c28e9a/large.pt",
|
||||
"large-v1": "https://openaipublic.azureedge.net/main/whisper/models/e4b87e7e0bf463eb8e6956e646f1e277e901512310def2c24bf0e11bd3c28e9a/large-v1.pt",
|
||||
"large-v2": "https://openaipublic.azureedge.net/main/whisper/models/81f7c96c852ee8fc832187b0132e569d6c3065a3252ed18e56effd0b6a73e524/large-v2.pt",
|
||||
"large-v3": "https://openaipublic.azureedge.net/main/whisper/models/e5b1a55b89c1367dacf97e3e19bfd829a01529dbfdeefa8caeb59b3f1b81dadb/large-v3.pt",
|
||||
}
|
||||
|
||||
|
||||
def _download(url: str, root: str, in_memory: bool) -> Union[bytes, str]:
|
||||
os.makedirs(root, exist_ok=True)
|
||||
|
||||
expected_sha256 = url.split("/")[-2]
|
||||
download_target = os.path.join(root, os.path.basename(url))
|
||||
|
||||
if os.path.exists(download_target) and not os.path.isfile(download_target):
|
||||
raise RuntimeError(f"{download_target} exists and is not a regular file")
|
||||
|
||||
if os.path.isfile(download_target):
|
||||
model_bytes = open(download_target, "rb").read()
|
||||
if hashlib.sha256(model_bytes).hexdigest() == expected_sha256:
|
||||
return model_bytes if in_memory else download_target
|
||||
else:
|
||||
warnings.warn(f"{download_target} exists, but the SHA256 checksum does not match; re-downloading the file")
|
||||
|
||||
with urllib.request.urlopen(url) as source, open(download_target, "wb") as output:
|
||||
with tqdm(
|
||||
total=int(source.info().get("Content-Length")), ncols=80, unit="iB", unit_scale=True, unit_divisor=1024
|
||||
) as loop:
|
||||
while True:
|
||||
buffer = source.read(8192)
|
||||
if not buffer:
|
||||
break
|
||||
|
||||
output.write(buffer)
|
||||
loop.update(len(buffer))
|
||||
|
||||
model_bytes = open(download_target, "rb").read()
|
||||
if hashlib.sha256(model_bytes).hexdigest() != expected_sha256:
|
||||
raise RuntimeError(
|
||||
"Model has been downloaded but the SHA256 checksum does not not match. Please retry loading the model."
|
||||
)
|
||||
|
||||
return model_bytes if in_memory else download_target
|
||||
|
||||
|
||||
def available_models() -> List[str]:
|
||||
"""Returns the names of available models"""
|
||||
return list(_MODELS.keys())
|
||||
|
||||
|
||||
def load_model(
|
||||
name: str, device: Optional[Union[str, torch.device]] = None, download_root: str = None, in_memory: bool = False
|
||||
) -> Whisper:
|
||||
"""
|
||||
Load a Whisper ASR model
|
||||
|
||||
Parameters
|
||||
----------
|
||||
name : str
|
||||
one of the official model names listed by `whisper.available_models()`, or
|
||||
path to a model checkpoint containing the model dimensions and the model state_dict.
|
||||
device : Union[str, torch.device]
|
||||
the PyTorch device to put the model into
|
||||
download_root: str
|
||||
path to download the model files; by default, it uses "~/.cache/whisper"
|
||||
in_memory: bool
|
||||
whether to preload the model weights into host memory
|
||||
|
||||
Returns
|
||||
-------
|
||||
model : Whisper
|
||||
The Whisper ASR model instance
|
||||
"""
|
||||
|
||||
if device is None:
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
if download_root is None:
|
||||
download_root = os.getenv("XDG_CACHE_HOME", os.path.join(os.path.expanduser("~"), ".cache", "whisper"))
|
||||
|
||||
if name in _MODELS:
|
||||
checkpoint_file = _download(_MODELS[name], download_root, in_memory)
|
||||
elif os.path.isfile(name):
|
||||
checkpoint_file = open(name, "rb").read() if in_memory else name
|
||||
else:
|
||||
raise RuntimeError(f"Model {name} not found; available models = {available_models()}")
|
||||
|
||||
with io.BytesIO(checkpoint_file) if in_memory else open(checkpoint_file, "rb") as fp:
|
||||
checkpoint = torch.load(fp, map_location=device, weights_only=True)
|
||||
del checkpoint_file
|
||||
|
||||
dims = ModelDimensions(**checkpoint["dims"])
|
||||
model = Whisper(dims)
|
||||
model.load_state_dict(checkpoint["model_state_dict"])
|
||||
|
||||
del checkpoint
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return model.to(device)
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from .transcribe import cli
|
||||
|
||||
|
||||
cli()
|
||||
from .transcribe import cli
|
||||
|
||||
|
||||
cli()
|
||||
|
||||
+50001
-50001
File diff suppressed because it is too large
Load Diff
@@ -1 +1 @@
|
||||
{"<|endoftext|>": 50257}
|
||||
{"<|endoftext|>": 50257}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
+125
-119
@@ -1,119 +1,125 @@
|
||||
import os
|
||||
from functools import lru_cache
|
||||
from typing import Union
|
||||
|
||||
import ffmpeg
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from .utils import exact_div
|
||||
|
||||
# hard-coded audio hyperparameters
|
||||
SAMPLE_RATE = 16000
|
||||
N_FFT = 400
|
||||
N_MELS = 80
|
||||
HOP_LENGTH = 160
|
||||
CHUNK_LENGTH = 30
|
||||
N_SAMPLES = CHUNK_LENGTH * SAMPLE_RATE # 480000: number of samples in a chunk
|
||||
N_FRAMES = exact_div(N_SAMPLES, HOP_LENGTH) # 3000: number of frames in a mel spectrogram input
|
||||
|
||||
|
||||
def load_audio(file: str, sr: int = SAMPLE_RATE):
|
||||
"""
|
||||
Load an audio file and resample to 16kHz
|
||||
"""
|
||||
try:
|
||||
# Usa subprocess invece di ffmpeg-python
|
||||
import subprocess
|
||||
import numpy as np
|
||||
|
||||
cmd = ['ffmpeg', '-i', file, '-f', 'f32le', '-ac', '1', '-ar', str(sr), '-']
|
||||
proc = subprocess.Popen(cmd, stdout=subprocess.PIPE, stderr=subprocess.PIPE)
|
||||
stdout, stderr = proc.communicate()
|
||||
|
||||
if proc.returncode != 0:
|
||||
raise RuntimeError(f"Failed to load audio: {stderr.decode()}")
|
||||
|
||||
audio = np.frombuffer(stdout, dtype=np.float32)
|
||||
return audio
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error loading audio: {str(e)}")
|
||||
return None
|
||||
|
||||
|
||||
def pad_or_trim(array, length: int = N_SAMPLES, *, axis: int = -1):
|
||||
"""
|
||||
Pad or trim the audio array to N_SAMPLES, as expected by the encoder.
|
||||
"""
|
||||
if torch.is_tensor(array):
|
||||
if array.shape[axis] > length:
|
||||
array = array.index_select(dim=axis, index=torch.arange(length))
|
||||
|
||||
if array.shape[axis] < length:
|
||||
pad_widths = [(0, 0)] * array.ndim
|
||||
pad_widths[axis] = (0, length - array.shape[axis])
|
||||
array = F.pad(array, [pad for sizes in pad_widths[::-1] for pad in sizes])
|
||||
else:
|
||||
if array.shape[axis] > length:
|
||||
array = array.take(indices=range(length), axis=axis)
|
||||
|
||||
if array.shape[axis] < length:
|
||||
pad_widths = [(0, 0)] * array.ndim
|
||||
pad_widths[axis] = (0, length - array.shape[axis])
|
||||
array = np.pad(array, pad_widths)
|
||||
|
||||
return array
|
||||
|
||||
|
||||
@lru_cache(maxsize=None)
|
||||
def mel_filters(device, n_mels: int = N_MELS) -> torch.Tensor:
|
||||
"""
|
||||
load the mel filterbank matrix for projecting STFT into a Mel spectrogram.
|
||||
Allows decoupling librosa dependency; saved using:
|
||||
|
||||
np.savez_compressed(
|
||||
"mel_filters.npz",
|
||||
mel_80=librosa.filters.mel(sr=16000, n_fft=400, n_mels=80),
|
||||
)
|
||||
"""
|
||||
assert n_mels == 80, f"Unsupported n_mels: {n_mels}"
|
||||
with np.load(os.path.join(os.path.dirname(__file__), "assets", "mel_filters.npz")) as f:
|
||||
return torch.from_numpy(f[f"mel_{n_mels}"]).to(device)
|
||||
|
||||
|
||||
def log_mel_spectrogram(audio: Union[str, np.ndarray, torch.Tensor], n_mels: int = N_MELS):
|
||||
"""
|
||||
Compute the log-Mel spectrogram of
|
||||
|
||||
Parameters
|
||||
----------
|
||||
audio: Union[str, np.ndarray, torch.Tensor], shape = (*)
|
||||
The path to audio or either a NumPy array or Tensor containing the audio waveform in 16 kHz
|
||||
|
||||
n_mels: int
|
||||
The number of Mel-frequency filters, only 80 is supported
|
||||
|
||||
Returns
|
||||
-------
|
||||
torch.Tensor, shape = (80, n_frames)
|
||||
A Tensor that contains the Mel spectrogram
|
||||
"""
|
||||
if not torch.is_tensor(audio):
|
||||
if isinstance(audio, str):
|
||||
audio = load_audio(audio)
|
||||
audio = torch.from_numpy(audio)
|
||||
|
||||
window = torch.hann_window(N_FFT).to(audio.device)
|
||||
stft = torch.stft(audio, N_FFT, HOP_LENGTH, window=window, return_complex=True)
|
||||
|
||||
magnitudes = stft[:, :-1].abs() ** 2
|
||||
|
||||
filters = mel_filters(audio.device, n_mels)
|
||||
mel_spec = filters @ magnitudes
|
||||
|
||||
log_spec = torch.clamp(mel_spec, min=1e-10).log10()
|
||||
log_spec = torch.maximum(log_spec, log_spec.max() - 8.0)
|
||||
log_spec = (log_spec + 4.0) / 4.0
|
||||
return log_spec
|
||||
import os
|
||||
from functools import lru_cache
|
||||
from typing import Union
|
||||
|
||||
import ffmpeg
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from .utils import exact_div
|
||||
|
||||
# hard-coded audio hyperparameters
|
||||
SAMPLE_RATE = 16000
|
||||
N_FFT = 400
|
||||
N_MELS = 80
|
||||
HOP_LENGTH = 160
|
||||
CHUNK_LENGTH = 30
|
||||
N_SAMPLES = CHUNK_LENGTH * SAMPLE_RATE # 480000: number of samples in a chunk
|
||||
N_FRAMES = exact_div(N_SAMPLES, HOP_LENGTH) # 3000: number of frames in a mel spectrogram input
|
||||
|
||||
|
||||
def load_audio(file: str, sr: int = SAMPLE_RATE):
|
||||
"""
|
||||
Open an audio file and read as mono waveform, resampling as necessary
|
||||
|
||||
Parameters
|
||||
----------
|
||||
file: str
|
||||
The audio file to open
|
||||
|
||||
sr: int
|
||||
The sample rate to resample the audio if necessary
|
||||
|
||||
Returns
|
||||
-------
|
||||
A NumPy array containing the audio waveform, in float32 dtype.
|
||||
"""
|
||||
try:
|
||||
# This launches a subprocess to decode audio while down-mixing and resampling as necessary.
|
||||
# Requires the ffmpeg CLI and `ffmpeg-python` package to be installed.
|
||||
out, _ = (
|
||||
ffmpeg.input(file, threads=0)
|
||||
.output("-", format="s16le", acodec="pcm_s16le", ac=1, ar=sr)
|
||||
.run(cmd=["ffmpeg", "-nostdin"], capture_stdout=True, capture_stderr=True)
|
||||
)
|
||||
except ffmpeg.Error as e:
|
||||
raise RuntimeError(f"Failed to load audio: {e.stderr.decode()}") from e
|
||||
|
||||
return np.frombuffer(out, np.int16).flatten().astype(np.float32) / 32768.0
|
||||
|
||||
|
||||
def pad_or_trim(array, length: int = N_SAMPLES, *, axis: int = -1):
|
||||
"""
|
||||
Pad or trim the audio array to N_SAMPLES, as expected by the encoder.
|
||||
"""
|
||||
if torch.is_tensor(array):
|
||||
if array.shape[axis] > length:
|
||||
array = array.index_select(dim=axis, index=torch.arange(length))
|
||||
|
||||
if array.shape[axis] < length:
|
||||
pad_widths = [(0, 0)] * array.ndim
|
||||
pad_widths[axis] = (0, length - array.shape[axis])
|
||||
array = F.pad(array, [pad for sizes in pad_widths[::-1] for pad in sizes])
|
||||
else:
|
||||
if array.shape[axis] > length:
|
||||
array = array.take(indices=range(length), axis=axis)
|
||||
|
||||
if array.shape[axis] < length:
|
||||
pad_widths = [(0, 0)] * array.ndim
|
||||
pad_widths[axis] = (0, length - array.shape[axis])
|
||||
array = np.pad(array, pad_widths)
|
||||
|
||||
return array
|
||||
|
||||
|
||||
@lru_cache(maxsize=None)
|
||||
def mel_filters(device, n_mels: int = N_MELS) -> torch.Tensor:
|
||||
"""
|
||||
load the mel filterbank matrix for projecting STFT into a Mel spectrogram.
|
||||
Allows decoupling librosa dependency; saved using:
|
||||
|
||||
np.savez_compressed(
|
||||
"mel_filters.npz",
|
||||
mel_80=librosa.filters.mel(sr=16000, n_fft=400, n_mels=80),
|
||||
)
|
||||
"""
|
||||
assert n_mels == 80, f"Unsupported n_mels: {n_mels}"
|
||||
with np.load(os.path.join(os.path.dirname(__file__), "assets", "mel_filters.npz")) as f:
|
||||
return torch.from_numpy(f[f"mel_{n_mels}"]).to(device)
|
||||
|
||||
|
||||
def log_mel_spectrogram(audio: Union[str, np.ndarray, torch.Tensor], n_mels: int = N_MELS):
|
||||
"""
|
||||
Compute the log-Mel spectrogram of
|
||||
|
||||
Parameters
|
||||
----------
|
||||
audio: Union[str, np.ndarray, torch.Tensor], shape = (*)
|
||||
The path to audio or either a NumPy array or Tensor containing the audio waveform in 16 kHz
|
||||
|
||||
n_mels: int
|
||||
The number of Mel-frequency filters, only 80 is supported
|
||||
|
||||
Returns
|
||||
-------
|
||||
torch.Tensor, shape = (80, n_frames)
|
||||
A Tensor that contains the Mel spectrogram
|
||||
"""
|
||||
if not torch.is_tensor(audio):
|
||||
if isinstance(audio, str):
|
||||
audio = load_audio(audio)
|
||||
audio = torch.from_numpy(audio)
|
||||
|
||||
window = torch.hann_window(N_FFT).to(audio.device)
|
||||
stft = torch.stft(audio, N_FFT, HOP_LENGTH, window=window, return_complex=True)
|
||||
|
||||
magnitudes = stft[:, :-1].abs() ** 2
|
||||
|
||||
filters = mel_filters(audio.device, n_mels)
|
||||
mel_spec = filters @ magnitudes
|
||||
|
||||
log_spec = torch.clamp(mel_spec, min=1e-10).log10()
|
||||
log_spec = torch.maximum(log_spec, log_spec.max() - 8.0)
|
||||
log_spec = (log_spec + 4.0) / 4.0
|
||||
return log_spec
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
+290
-290
@@ -1,290 +1,290 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict
|
||||
from typing import Iterable, Optional
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
from torch import nn
|
||||
|
||||
from .transcribe import transcribe as transcribe_function
|
||||
from .decoding import detect_language as detect_language_function, decode as decode_function
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelDimensions:
|
||||
n_mels: int
|
||||
n_audio_ctx: int
|
||||
n_audio_state: int
|
||||
n_audio_head: int
|
||||
n_audio_layer: int
|
||||
n_vocab: int
|
||||
n_text_ctx: int
|
||||
n_text_state: int
|
||||
n_text_head: int
|
||||
n_text_layer: int
|
||||
|
||||
|
||||
class LayerNorm(nn.LayerNorm):
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
return super().forward(x.float()).type(x.dtype)
|
||||
|
||||
|
||||
class Linear(nn.Linear):
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
return F.linear(
|
||||
x, self.weight.to(x.dtype), None if self.bias is None else self.bias.to(x.dtype)
|
||||
)
|
||||
|
||||
|
||||
class Conv1d(nn.Conv1d):
|
||||
def _conv_forward(self, x: Tensor, weight: Tensor, bias: Optional[Tensor]) -> Tensor:
|
||||
return super()._conv_forward(
|
||||
x, weight.to(x.dtype), None if bias is None else bias.to(x.dtype)
|
||||
)
|
||||
|
||||
|
||||
def sinusoids(length, channels, max_timescale=10000):
|
||||
"""Returns sinusoids for positional embedding"""
|
||||
assert channels % 2 == 0
|
||||
log_timescale_increment = np.log(max_timescale) / (channels // 2 - 1)
|
||||
inv_timescales = torch.exp(-log_timescale_increment * torch.arange(channels // 2))
|
||||
scaled_time = torch.arange(length)[:, np.newaxis] * inv_timescales[np.newaxis, :]
|
||||
return torch.cat([torch.sin(scaled_time), torch.cos(scaled_time)], dim=1)
|
||||
|
||||
|
||||
class MultiHeadAttention(nn.Module):
|
||||
def __init__(self, n_state: int, n_head: int):
|
||||
super().__init__()
|
||||
self.n_head = n_head
|
||||
self.query = Linear(n_state, n_state)
|
||||
self.key = Linear(n_state, n_state, bias=False)
|
||||
self.value = Linear(n_state, n_state)
|
||||
self.out = Linear(n_state, n_state)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: Tensor,
|
||||
xa: Optional[Tensor] = None,
|
||||
mask: Optional[Tensor] = None,
|
||||
kv_cache: Optional[dict] = None,
|
||||
):
|
||||
q = self.query(x)
|
||||
|
||||
if kv_cache is None or xa is None:
|
||||
# hooks, if installed (i.e. kv_cache is not None), will prepend the cached kv tensors;
|
||||
# otherwise, perform key/value projections for self- or cross-attention as usual.
|
||||
k = self.key(x if xa is None else xa)
|
||||
v = self.value(x if xa is None else xa)
|
||||
else:
|
||||
# for cross-attention, calculate keys and values once and reuse in subsequent calls.
|
||||
k = kv_cache.get(self.key, self.key(xa))
|
||||
v = kv_cache.get(self.value, self.value(xa))
|
||||
|
||||
wv = self.qkv_attention(q, k, v, mask)
|
||||
return self.out(wv)
|
||||
|
||||
def qkv_attention(self, q: Tensor, k: Tensor, v: Tensor, mask: Optional[Tensor] = None):
|
||||
n_batch, n_ctx, n_state = q.shape
|
||||
scale = (n_state // self.n_head) ** -0.25
|
||||
q = q.view(*q.shape[:2], self.n_head, -1).permute(0, 2, 1, 3) * scale
|
||||
k = k.view(*k.shape[:2], self.n_head, -1).permute(0, 2, 3, 1) * scale
|
||||
v = v.view(*v.shape[:2], self.n_head, -1).permute(0, 2, 1, 3)
|
||||
|
||||
qk = q @ k
|
||||
if mask is not None:
|
||||
qk = qk + mask[:n_ctx, :n_ctx]
|
||||
|
||||
w = F.softmax(qk.float(), dim=-1).to(q.dtype)
|
||||
return (w @ v).permute(0, 2, 1, 3).flatten(start_dim=2)
|
||||
|
||||
|
||||
class ResidualAttentionBlock(nn.Module):
|
||||
def __init__(self, n_state: int, n_head: int, cross_attention: bool = False):
|
||||
super().__init__()
|
||||
|
||||
self.attn = MultiHeadAttention(n_state, n_head)
|
||||
self.attn_ln = LayerNorm(n_state)
|
||||
|
||||
self.cross_attn = MultiHeadAttention(n_state, n_head) if cross_attention else None
|
||||
self.cross_attn_ln = LayerNorm(n_state) if cross_attention else None
|
||||
|
||||
n_mlp = n_state * 4
|
||||
self.mlp = nn.Sequential(Linear(n_state, n_mlp), nn.GELU(), Linear(n_mlp, n_state))
|
||||
self.mlp_ln = LayerNorm(n_state)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: Tensor,
|
||||
xa: Optional[Tensor] = None,
|
||||
mask: Optional[Tensor] = None,
|
||||
kv_cache: Optional[dict] = None,
|
||||
):
|
||||
x = x + self.attn(self.attn_ln(x), mask=mask, kv_cache=kv_cache)
|
||||
if self.cross_attn:
|
||||
x = x + self.cross_attn(self.cross_attn_ln(x), xa, kv_cache=kv_cache)
|
||||
x = x + self.mlp(self.mlp_ln(x))
|
||||
return x
|
||||
|
||||
|
||||
class AudioEncoder(nn.Module):
|
||||
def __init__(self, n_mels: int, n_ctx: int, n_state: int, n_head: int, n_layer: int):
|
||||
super().__init__()
|
||||
self.conv1 = Conv1d(n_mels, n_state, kernel_size=3, padding=1)
|
||||
self.conv2 = Conv1d(n_state, n_state, kernel_size=3, stride=2, padding=1)
|
||||
self.register_buffer("positional_embedding", sinusoids(n_ctx, n_state))
|
||||
|
||||
self.blocks: Iterable[ResidualAttentionBlock] = nn.ModuleList(
|
||||
[ResidualAttentionBlock(n_state, n_head) for _ in range(n_layer)]
|
||||
)
|
||||
self.ln_post = LayerNorm(n_state)
|
||||
|
||||
def forward(self, x: Tensor, include_embeddings: bool = False):
|
||||
"""
|
||||
x : torch.Tensor, shape = (batch_size, n_mels, n_ctx)
|
||||
the mel spectrogram of the audio
|
||||
include_embeddings: bool
|
||||
whether to include intermediate steps in the output
|
||||
"""
|
||||
x = F.gelu(self.conv1(x))
|
||||
x = F.gelu(self.conv2(x))
|
||||
x = x.permute(0, 2, 1)
|
||||
|
||||
assert x.shape[1:] == self.positional_embedding.shape, "incorrect audio shape"
|
||||
x = (x + self.positional_embedding).to(x.dtype)
|
||||
|
||||
if include_embeddings:
|
||||
embeddings = [x.cpu().detach().numpy()]
|
||||
|
||||
for block in self.blocks:
|
||||
x = block(x)
|
||||
if include_embeddings:
|
||||
embeddings.append(x.cpu().detach().numpy())
|
||||
|
||||
x = self.ln_post(x)
|
||||
|
||||
if include_embeddings:
|
||||
embeddings = np.stack(embeddings, axis=1)
|
||||
return x, embeddings
|
||||
else:
|
||||
return x
|
||||
|
||||
|
||||
class TextDecoder(nn.Module):
|
||||
def __init__(self, n_vocab: int, n_ctx: int, n_state: int, n_head: int, n_layer: int):
|
||||
super().__init__()
|
||||
|
||||
self.token_embedding = nn.Embedding(n_vocab, n_state)
|
||||
self.positional_embedding = nn.Parameter(torch.empty(n_ctx, n_state))
|
||||
|
||||
self.blocks: Iterable[ResidualAttentionBlock] = nn.ModuleList(
|
||||
[ResidualAttentionBlock(n_state, n_head, cross_attention=True) for _ in range(n_layer)]
|
||||
)
|
||||
self.ln = LayerNorm(n_state)
|
||||
|
||||
mask = torch.empty(n_ctx, n_ctx).fill_(-np.inf).triu_(1)
|
||||
self.register_buffer("mask", mask, persistent=False)
|
||||
|
||||
def forward(self, x: Tensor, xa: Tensor, kv_cache: Optional[dict] = None, include_embeddings: bool = False):
|
||||
"""
|
||||
x : torch.LongTensor, shape = (batch_size, <= n_ctx)
|
||||
the text tokens
|
||||
xa : torch.Tensor, shape = (batch_size, n_mels, n_audio_ctx)
|
||||
the encoded audio features to be attended on
|
||||
include_embeddings : bool
|
||||
Whether to include intermediate values in the output to this function
|
||||
"""
|
||||
offset = next(iter(kv_cache.values())).shape[1] if kv_cache else 0
|
||||
x = self.token_embedding(x) + self.positional_embedding[offset : offset + x.shape[-1]]
|
||||
x = x.to(xa.dtype)
|
||||
|
||||
if include_embeddings:
|
||||
embeddings = [x.cpu().detach().numpy()]
|
||||
|
||||
for block in self.blocks:
|
||||
x = block(x, xa, mask=self.mask, kv_cache=kv_cache)
|
||||
if include_embeddings:
|
||||
embeddings.append(x.cpu().detach().numpy())
|
||||
|
||||
x = self.ln(x)
|
||||
logits = (x @ torch.transpose(self.token_embedding.weight.to(x.dtype), 0, 1)).float()
|
||||
|
||||
if include_embeddings:
|
||||
embeddings = np.stack(embeddings, axis=1)
|
||||
return logits, embeddings
|
||||
else:
|
||||
return logits
|
||||
|
||||
|
||||
class Whisper(nn.Module):
|
||||
def __init__(self, dims: ModelDimensions):
|
||||
super().__init__()
|
||||
self.dims = dims
|
||||
self.encoder = AudioEncoder(
|
||||
self.dims.n_mels,
|
||||
self.dims.n_audio_ctx,
|
||||
self.dims.n_audio_state,
|
||||
self.dims.n_audio_head,
|
||||
self.dims.n_audio_layer,
|
||||
)
|
||||
self.decoder = TextDecoder(
|
||||
self.dims.n_vocab,
|
||||
self.dims.n_text_ctx,
|
||||
self.dims.n_text_state,
|
||||
self.dims.n_text_head,
|
||||
self.dims.n_text_layer,
|
||||
)
|
||||
|
||||
def embed_audio(self, mel: torch.Tensor):
|
||||
return self.encoder.forward(mel)
|
||||
|
||||
def logits(self, tokens: torch.Tensor, audio_features: torch.Tensor):
|
||||
return self.decoder.forward(tokens, audio_features)
|
||||
|
||||
def forward(self, mel: torch.Tensor, tokens: torch.Tensor) -> Dict[str, torch.Tensor]:
|
||||
return self.decoder(tokens, self.encoder(mel))
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
return next(self.parameters()).device
|
||||
|
||||
@property
|
||||
def is_multilingual(self):
|
||||
return self.dims.n_vocab == 51865
|
||||
|
||||
def install_kv_cache_hooks(self, cache: Optional[dict] = None):
|
||||
"""
|
||||
The `MultiHeadAttention` module optionally accepts `kv_cache` which stores the key and value
|
||||
tensors calculated for the previous positions. This method returns a dictionary that stores
|
||||
all caches, and the necessary hooks for the key and value projection modules that save the
|
||||
intermediate tensors to be reused during later calculations.
|
||||
|
||||
Returns
|
||||
-------
|
||||
cache : Dict[nn.Module, torch.Tensor]
|
||||
A dictionary object mapping the key/value projection modules to its cache
|
||||
hooks : List[RemovableHandle]
|
||||
List of PyTorch RemovableHandle objects to stop the hooks to be called
|
||||
"""
|
||||
cache = {**cache} if cache is not None else {}
|
||||
hooks = []
|
||||
|
||||
def save_to_cache(module, _, output):
|
||||
if module not in cache or output.shape[1] > self.decoder.positional_embedding.shape[0]:
|
||||
cache[module] = output # save as-is, for the first token or cross attention
|
||||
else:
|
||||
cache[module] = torch.cat([cache[module], output], dim=1).detach()
|
||||
return cache[module]
|
||||
|
||||
def install_hooks(layer: nn.Module):
|
||||
if isinstance(layer, MultiHeadAttention):
|
||||
hooks.append(layer.key.register_forward_hook(save_to_cache))
|
||||
hooks.append(layer.value.register_forward_hook(save_to_cache))
|
||||
|
||||
self.decoder.apply(install_hooks)
|
||||
return cache, hooks
|
||||
|
||||
detect_language = detect_language_function
|
||||
transcribe = transcribe_function
|
||||
decode = decode_function
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict
|
||||
from typing import Iterable, Optional
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
from torch import nn
|
||||
|
||||
from .transcribe import transcribe as transcribe_function
|
||||
from .decoding import detect_language as detect_language_function, decode as decode_function
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelDimensions:
|
||||
n_mels: int
|
||||
n_audio_ctx: int
|
||||
n_audio_state: int
|
||||
n_audio_head: int
|
||||
n_audio_layer: int
|
||||
n_vocab: int
|
||||
n_text_ctx: int
|
||||
n_text_state: int
|
||||
n_text_head: int
|
||||
n_text_layer: int
|
||||
|
||||
|
||||
class LayerNorm(nn.LayerNorm):
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
return super().forward(x.float()).type(x.dtype)
|
||||
|
||||
|
||||
class Linear(nn.Linear):
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
return F.linear(
|
||||
x, self.weight.to(x.dtype), None if self.bias is None else self.bias.to(x.dtype)
|
||||
)
|
||||
|
||||
|
||||
class Conv1d(nn.Conv1d):
|
||||
def _conv_forward(self, x: Tensor, weight: Tensor, bias: Optional[Tensor]) -> Tensor:
|
||||
return super()._conv_forward(
|
||||
x, weight.to(x.dtype), None if bias is None else bias.to(x.dtype)
|
||||
)
|
||||
|
||||
|
||||
def sinusoids(length, channels, max_timescale=10000):
|
||||
"""Returns sinusoids for positional embedding"""
|
||||
assert channels % 2 == 0
|
||||
log_timescale_increment = np.log(max_timescale) / (channels // 2 - 1)
|
||||
inv_timescales = torch.exp(-log_timescale_increment * torch.arange(channels // 2))
|
||||
scaled_time = torch.arange(length)[:, np.newaxis] * inv_timescales[np.newaxis, :]
|
||||
return torch.cat([torch.sin(scaled_time), torch.cos(scaled_time)], dim=1)
|
||||
|
||||
|
||||
class MultiHeadAttention(nn.Module):
|
||||
def __init__(self, n_state: int, n_head: int):
|
||||
super().__init__()
|
||||
self.n_head = n_head
|
||||
self.query = Linear(n_state, n_state)
|
||||
self.key = Linear(n_state, n_state, bias=False)
|
||||
self.value = Linear(n_state, n_state)
|
||||
self.out = Linear(n_state, n_state)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: Tensor,
|
||||
xa: Optional[Tensor] = None,
|
||||
mask: Optional[Tensor] = None,
|
||||
kv_cache: Optional[dict] = None,
|
||||
):
|
||||
q = self.query(x)
|
||||
|
||||
if kv_cache is None or xa is None:
|
||||
# hooks, if installed (i.e. kv_cache is not None), will prepend the cached kv tensors;
|
||||
# otherwise, perform key/value projections for self- or cross-attention as usual.
|
||||
k = self.key(x if xa is None else xa)
|
||||
v = self.value(x if xa is None else xa)
|
||||
else:
|
||||
# for cross-attention, calculate keys and values once and reuse in subsequent calls.
|
||||
k = kv_cache.get(self.key, self.key(xa))
|
||||
v = kv_cache.get(self.value, self.value(xa))
|
||||
|
||||
wv = self.qkv_attention(q, k, v, mask)
|
||||
return self.out(wv)
|
||||
|
||||
def qkv_attention(self, q: Tensor, k: Tensor, v: Tensor, mask: Optional[Tensor] = None):
|
||||
n_batch, n_ctx, n_state = q.shape
|
||||
scale = (n_state // self.n_head) ** -0.25
|
||||
q = q.view(*q.shape[:2], self.n_head, -1).permute(0, 2, 1, 3) * scale
|
||||
k = k.view(*k.shape[:2], self.n_head, -1).permute(0, 2, 3, 1) * scale
|
||||
v = v.view(*v.shape[:2], self.n_head, -1).permute(0, 2, 1, 3)
|
||||
|
||||
qk = q @ k
|
||||
if mask is not None:
|
||||
qk = qk + mask[:n_ctx, :n_ctx]
|
||||
|
||||
w = F.softmax(qk.float(), dim=-1).to(q.dtype)
|
||||
return (w @ v).permute(0, 2, 1, 3).flatten(start_dim=2)
|
||||
|
||||
|
||||
class ResidualAttentionBlock(nn.Module):
|
||||
def __init__(self, n_state: int, n_head: int, cross_attention: bool = False):
|
||||
super().__init__()
|
||||
|
||||
self.attn = MultiHeadAttention(n_state, n_head)
|
||||
self.attn_ln = LayerNorm(n_state)
|
||||
|
||||
self.cross_attn = MultiHeadAttention(n_state, n_head) if cross_attention else None
|
||||
self.cross_attn_ln = LayerNorm(n_state) if cross_attention else None
|
||||
|
||||
n_mlp = n_state * 4
|
||||
self.mlp = nn.Sequential(Linear(n_state, n_mlp), nn.GELU(), Linear(n_mlp, n_state))
|
||||
self.mlp_ln = LayerNorm(n_state)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: Tensor,
|
||||
xa: Optional[Tensor] = None,
|
||||
mask: Optional[Tensor] = None,
|
||||
kv_cache: Optional[dict] = None,
|
||||
):
|
||||
x = x + self.attn(self.attn_ln(x), mask=mask, kv_cache=kv_cache)
|
||||
if self.cross_attn:
|
||||
x = x + self.cross_attn(self.cross_attn_ln(x), xa, kv_cache=kv_cache)
|
||||
x = x + self.mlp(self.mlp_ln(x))
|
||||
return x
|
||||
|
||||
|
||||
class AudioEncoder(nn.Module):
|
||||
def __init__(self, n_mels: int, n_ctx: int, n_state: int, n_head: int, n_layer: int):
|
||||
super().__init__()
|
||||
self.conv1 = Conv1d(n_mels, n_state, kernel_size=3, padding=1)
|
||||
self.conv2 = Conv1d(n_state, n_state, kernel_size=3, stride=2, padding=1)
|
||||
self.register_buffer("positional_embedding", sinusoids(n_ctx, n_state))
|
||||
|
||||
self.blocks: Iterable[ResidualAttentionBlock] = nn.ModuleList(
|
||||
[ResidualAttentionBlock(n_state, n_head) for _ in range(n_layer)]
|
||||
)
|
||||
self.ln_post = LayerNorm(n_state)
|
||||
|
||||
def forward(self, x: Tensor, include_embeddings: bool = False):
|
||||
"""
|
||||
x : torch.Tensor, shape = (batch_size, n_mels, n_ctx)
|
||||
the mel spectrogram of the audio
|
||||
include_embeddings: bool
|
||||
whether to include intermediate steps in the output
|
||||
"""
|
||||
x = F.gelu(self.conv1(x))
|
||||
x = F.gelu(self.conv2(x))
|
||||
x = x.permute(0, 2, 1)
|
||||
|
||||
assert x.shape[1:] == self.positional_embedding.shape, "incorrect audio shape"
|
||||
x = (x + self.positional_embedding).to(x.dtype)
|
||||
|
||||
if include_embeddings:
|
||||
embeddings = [x.cpu().detach().numpy()]
|
||||
|
||||
for block in self.blocks:
|
||||
x = block(x)
|
||||
if include_embeddings:
|
||||
embeddings.append(x.cpu().detach().numpy())
|
||||
|
||||
x = self.ln_post(x)
|
||||
|
||||
if include_embeddings:
|
||||
embeddings = np.stack(embeddings, axis=1)
|
||||
return x, embeddings
|
||||
else:
|
||||
return x
|
||||
|
||||
|
||||
class TextDecoder(nn.Module):
|
||||
def __init__(self, n_vocab: int, n_ctx: int, n_state: int, n_head: int, n_layer: int):
|
||||
super().__init__()
|
||||
|
||||
self.token_embedding = nn.Embedding(n_vocab, n_state)
|
||||
self.positional_embedding = nn.Parameter(torch.empty(n_ctx, n_state))
|
||||
|
||||
self.blocks: Iterable[ResidualAttentionBlock] = nn.ModuleList(
|
||||
[ResidualAttentionBlock(n_state, n_head, cross_attention=True) for _ in range(n_layer)]
|
||||
)
|
||||
self.ln = LayerNorm(n_state)
|
||||
|
||||
mask = torch.empty(n_ctx, n_ctx).fill_(-np.inf).triu_(1)
|
||||
self.register_buffer("mask", mask, persistent=False)
|
||||
|
||||
def forward(self, x: Tensor, xa: Tensor, kv_cache: Optional[dict] = None, include_embeddings: bool = False):
|
||||
"""
|
||||
x : torch.LongTensor, shape = (batch_size, <= n_ctx)
|
||||
the text tokens
|
||||
xa : torch.Tensor, shape = (batch_size, n_mels, n_audio_ctx)
|
||||
the encoded audio features to be attended on
|
||||
include_embeddings : bool
|
||||
Whether to include intermediate values in the output to this function
|
||||
"""
|
||||
offset = next(iter(kv_cache.values())).shape[1] if kv_cache else 0
|
||||
x = self.token_embedding(x) + self.positional_embedding[offset : offset + x.shape[-1]]
|
||||
x = x.to(xa.dtype)
|
||||
|
||||
if include_embeddings:
|
||||
embeddings = [x.cpu().detach().numpy()]
|
||||
|
||||
for block in self.blocks:
|
||||
x = block(x, xa, mask=self.mask, kv_cache=kv_cache)
|
||||
if include_embeddings:
|
||||
embeddings.append(x.cpu().detach().numpy())
|
||||
|
||||
x = self.ln(x)
|
||||
logits = (x @ torch.transpose(self.token_embedding.weight.to(x.dtype), 0, 1)).float()
|
||||
|
||||
if include_embeddings:
|
||||
embeddings = np.stack(embeddings, axis=1)
|
||||
return logits, embeddings
|
||||
else:
|
||||
return logits
|
||||
|
||||
|
||||
class Whisper(nn.Module):
|
||||
def __init__(self, dims: ModelDimensions):
|
||||
super().__init__()
|
||||
self.dims = dims
|
||||
self.encoder = AudioEncoder(
|
||||
self.dims.n_mels,
|
||||
self.dims.n_audio_ctx,
|
||||
self.dims.n_audio_state,
|
||||
self.dims.n_audio_head,
|
||||
self.dims.n_audio_layer,
|
||||
)
|
||||
self.decoder = TextDecoder(
|
||||
self.dims.n_vocab,
|
||||
self.dims.n_text_ctx,
|
||||
self.dims.n_text_state,
|
||||
self.dims.n_text_head,
|
||||
self.dims.n_text_layer,
|
||||
)
|
||||
|
||||
def embed_audio(self, mel: torch.Tensor):
|
||||
return self.encoder.forward(mel)
|
||||
|
||||
def logits(self, tokens: torch.Tensor, audio_features: torch.Tensor):
|
||||
return self.decoder.forward(tokens, audio_features)
|
||||
|
||||
def forward(self, mel: torch.Tensor, tokens: torch.Tensor) -> Dict[str, torch.Tensor]:
|
||||
return self.decoder(tokens, self.encoder(mel))
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
return next(self.parameters()).device
|
||||
|
||||
@property
|
||||
def is_multilingual(self):
|
||||
return self.dims.n_vocab == 51865
|
||||
|
||||
def install_kv_cache_hooks(self, cache: Optional[dict] = None):
|
||||
"""
|
||||
The `MultiHeadAttention` module optionally accepts `kv_cache` which stores the key and value
|
||||
tensors calculated for the previous positions. This method returns a dictionary that stores
|
||||
all caches, and the necessary hooks for the key and value projection modules that save the
|
||||
intermediate tensors to be reused during later calculations.
|
||||
|
||||
Returns
|
||||
-------
|
||||
cache : Dict[nn.Module, torch.Tensor]
|
||||
A dictionary object mapping the key/value projection modules to its cache
|
||||
hooks : List[RemovableHandle]
|
||||
List of PyTorch RemovableHandle objects to stop the hooks to be called
|
||||
"""
|
||||
cache = {**cache} if cache is not None else {}
|
||||
hooks = []
|
||||
|
||||
def save_to_cache(module, _, output):
|
||||
if module not in cache or output.shape[1] > self.decoder.positional_embedding.shape[0]:
|
||||
cache[module] = output # save as-is, for the first token or cross attention
|
||||
else:
|
||||
cache[module] = torch.cat([cache[module], output], dim=1).detach()
|
||||
return cache[module]
|
||||
|
||||
def install_hooks(layer: nn.Module):
|
||||
if isinstance(layer, MultiHeadAttention):
|
||||
hooks.append(layer.key.register_forward_hook(save_to_cache))
|
||||
hooks.append(layer.value.register_forward_hook(save_to_cache))
|
||||
|
||||
self.decoder.apply(install_hooks)
|
||||
return cache, hooks
|
||||
|
||||
detect_language = detect_language_function
|
||||
transcribe = transcribe_function
|
||||
decode = decode_function
|
||||
|
||||
@@ -1,2 +1,2 @@
|
||||
from .basic import BasicTextNormalizer
|
||||
from .english import EnglishTextNormalizer
|
||||
from .basic import BasicTextNormalizer
|
||||
from .english import EnglishTextNormalizer
|
||||
|
||||
@@ -1,71 +1,71 @@
|
||||
import re
|
||||
import unicodedata
|
||||
|
||||
import regex
|
||||
|
||||
# non-ASCII letters that are not separated by "NFKD" normalization
|
||||
ADDITIONAL_DIACRITICS = {
|
||||
"œ": "oe",
|
||||
"Œ": "OE",
|
||||
"ø": "o",
|
||||
"Ø": "O",
|
||||
"æ": "ae",
|
||||
"Æ": "AE",
|
||||
"ß": "ss",
|
||||
"ẞ": "SS",
|
||||
"đ": "d",
|
||||
"Đ": "D",
|
||||
"ð": "d",
|
||||
"Ð": "D",
|
||||
"þ": "th",
|
||||
"Þ": "th",
|
||||
"ł": "l",
|
||||
"Ł": "L",
|
||||
}
|
||||
|
||||
|
||||
def remove_symbols_and_diacritics(s: str, keep=""):
|
||||
"""
|
||||
Replace any other markers, symbols, and punctuations with a space,
|
||||
and drop any diacritics (category 'Mn' and some manual mappings)
|
||||
"""
|
||||
return "".join(
|
||||
c
|
||||
if c in keep
|
||||
else ADDITIONAL_DIACRITICS[c]
|
||||
if c in ADDITIONAL_DIACRITICS
|
||||
else ""
|
||||
if unicodedata.category(c) == "Mn"
|
||||
else " "
|
||||
if unicodedata.category(c)[0] in "MSP"
|
||||
else c
|
||||
for c in unicodedata.normalize("NFKD", s)
|
||||
)
|
||||
|
||||
|
||||
def remove_symbols(s: str):
|
||||
"""
|
||||
Replace any other markers, symbols, punctuations with a space, keeping diacritics
|
||||
"""
|
||||
return "".join(
|
||||
" " if unicodedata.category(c)[0] in "MSP" else c for c in unicodedata.normalize("NFKC", s)
|
||||
)
|
||||
|
||||
|
||||
class BasicTextNormalizer:
|
||||
def __init__(self, remove_diacritics: bool = False, split_letters: bool = False):
|
||||
self.clean = remove_symbols_and_diacritics if remove_diacritics else remove_symbols
|
||||
self.split_letters = split_letters
|
||||
|
||||
def __call__(self, s: str):
|
||||
s = s.lower()
|
||||
s = re.sub(r"[<\[][^>\]]*[>\]]", "", s) # remove words between brackets
|
||||
s = re.sub(r"\(([^)]+?)\)", "", s) # remove words between parenthesis
|
||||
s = self.clean(s).lower()
|
||||
|
||||
if self.split_letters:
|
||||
s = " ".join(regex.findall(r"\X", s, regex.U))
|
||||
|
||||
s = re.sub(r"\s+", " ", s) # replace any successive whitespace characters with a space
|
||||
|
||||
return s
|
||||
import re
|
||||
import unicodedata
|
||||
|
||||
import regex
|
||||
|
||||
# non-ASCII letters that are not separated by "NFKD" normalization
|
||||
ADDITIONAL_DIACRITICS = {
|
||||
"œ": "oe",
|
||||
"Œ": "OE",
|
||||
"ø": "o",
|
||||
"Ø": "O",
|
||||
"æ": "ae",
|
||||
"Æ": "AE",
|
||||
"ß": "ss",
|
||||
"ẞ": "SS",
|
||||
"đ": "d",
|
||||
"Đ": "D",
|
||||
"ð": "d",
|
||||
"Ð": "D",
|
||||
"þ": "th",
|
||||
"Þ": "th",
|
||||
"ł": "l",
|
||||
"Ł": "L",
|
||||
}
|
||||
|
||||
|
||||
def remove_symbols_and_diacritics(s: str, keep=""):
|
||||
"""
|
||||
Replace any other markers, symbols, and punctuations with a space,
|
||||
and drop any diacritics (category 'Mn' and some manual mappings)
|
||||
"""
|
||||
return "".join(
|
||||
c
|
||||
if c in keep
|
||||
else ADDITIONAL_DIACRITICS[c]
|
||||
if c in ADDITIONAL_DIACRITICS
|
||||
else ""
|
||||
if unicodedata.category(c) == "Mn"
|
||||
else " "
|
||||
if unicodedata.category(c)[0] in "MSP"
|
||||
else c
|
||||
for c in unicodedata.normalize("NFKD", s)
|
||||
)
|
||||
|
||||
|
||||
def remove_symbols(s: str):
|
||||
"""
|
||||
Replace any other markers, symbols, punctuations with a space, keeping diacritics
|
||||
"""
|
||||
return "".join(
|
||||
" " if unicodedata.category(c)[0] in "MSP" else c for c in unicodedata.normalize("NFKC", s)
|
||||
)
|
||||
|
||||
|
||||
class BasicTextNormalizer:
|
||||
def __init__(self, remove_diacritics: bool = False, split_letters: bool = False):
|
||||
self.clean = remove_symbols_and_diacritics if remove_diacritics else remove_symbols
|
||||
self.split_letters = split_letters
|
||||
|
||||
def __call__(self, s: str):
|
||||
s = s.lower()
|
||||
s = re.sub(r"[<\[][^>\]]*[>\]]", "", s) # remove words between brackets
|
||||
s = re.sub(r"\(([^)]+?)\)", "", s) # remove words between parenthesis
|
||||
s = self.clean(s).lower()
|
||||
|
||||
if self.split_letters:
|
||||
s = " ".join(regex.findall(r"\X", s, regex.U))
|
||||
|
||||
s = re.sub(r"\s+", " ", s) # replace any successive whitespace characters with a space
|
||||
|
||||
return s
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -1,331 +1,331 @@
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from functools import lru_cache
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from transformers import GPT2TokenizerFast
|
||||
|
||||
LANGUAGES = {
|
||||
"en": "english",
|
||||
"zh": "chinese",
|
||||
"de": "german",
|
||||
"es": "spanish",
|
||||
"ru": "russian",
|
||||
"ko": "korean",
|
||||
"fr": "french",
|
||||
"ja": "japanese",
|
||||
"pt": "portuguese",
|
||||
"tr": "turkish",
|
||||
"pl": "polish",
|
||||
"ca": "catalan",
|
||||
"nl": "dutch",
|
||||
"ar": "arabic",
|
||||
"sv": "swedish",
|
||||
"it": "italian",
|
||||
"id": "indonesian",
|
||||
"hi": "hindi",
|
||||
"fi": "finnish",
|
||||
"vi": "vietnamese",
|
||||
"iw": "hebrew",
|
||||
"uk": "ukrainian",
|
||||
"el": "greek",
|
||||
"ms": "malay",
|
||||
"cs": "czech",
|
||||
"ro": "romanian",
|
||||
"da": "danish",
|
||||
"hu": "hungarian",
|
||||
"ta": "tamil",
|
||||
"no": "norwegian",
|
||||
"th": "thai",
|
||||
"ur": "urdu",
|
||||
"hr": "croatian",
|
||||
"bg": "bulgarian",
|
||||
"lt": "lithuanian",
|
||||
"la": "latin",
|
||||
"mi": "maori",
|
||||
"ml": "malayalam",
|
||||
"cy": "welsh",
|
||||
"sk": "slovak",
|
||||
"te": "telugu",
|
||||
"fa": "persian",
|
||||
"lv": "latvian",
|
||||
"bn": "bengali",
|
||||
"sr": "serbian",
|
||||
"az": "azerbaijani",
|
||||
"sl": "slovenian",
|
||||
"kn": "kannada",
|
||||
"et": "estonian",
|
||||
"mk": "macedonian",
|
||||
"br": "breton",
|
||||
"eu": "basque",
|
||||
"is": "icelandic",
|
||||
"hy": "armenian",
|
||||
"ne": "nepali",
|
||||
"mn": "mongolian",
|
||||
"bs": "bosnian",
|
||||
"kk": "kazakh",
|
||||
"sq": "albanian",
|
||||
"sw": "swahili",
|
||||
"gl": "galician",
|
||||
"mr": "marathi",
|
||||
"pa": "punjabi",
|
||||
"si": "sinhala",
|
||||
"km": "khmer",
|
||||
"sn": "shona",
|
||||
"yo": "yoruba",
|
||||
"so": "somali",
|
||||
"af": "afrikaans",
|
||||
"oc": "occitan",
|
||||
"ka": "georgian",
|
||||
"be": "belarusian",
|
||||
"tg": "tajik",
|
||||
"sd": "sindhi",
|
||||
"gu": "gujarati",
|
||||
"am": "amharic",
|
||||
"yi": "yiddish",
|
||||
"lo": "lao",
|
||||
"uz": "uzbek",
|
||||
"fo": "faroese",
|
||||
"ht": "haitian creole",
|
||||
"ps": "pashto",
|
||||
"tk": "turkmen",
|
||||
"nn": "nynorsk",
|
||||
"mt": "maltese",
|
||||
"sa": "sanskrit",
|
||||
"lb": "luxembourgish",
|
||||
"my": "myanmar",
|
||||
"bo": "tibetan",
|
||||
"tl": "tagalog",
|
||||
"mg": "malagasy",
|
||||
"as": "assamese",
|
||||
"tt": "tatar",
|
||||
"haw": "hawaiian",
|
||||
"ln": "lingala",
|
||||
"ha": "hausa",
|
||||
"ba": "bashkir",
|
||||
"jw": "javanese",
|
||||
"su": "sundanese",
|
||||
}
|
||||
|
||||
# language code lookup by name, with a few language aliases
|
||||
TO_LANGUAGE_CODE = {
|
||||
**{language: code for code, language in LANGUAGES.items()},
|
||||
"burmese": "my",
|
||||
"valencian": "ca",
|
||||
"flemish": "nl",
|
||||
"haitian": "ht",
|
||||
"letzeburgesch": "lb",
|
||||
"pushto": "ps",
|
||||
"panjabi": "pa",
|
||||
"moldavian": "ro",
|
||||
"moldovan": "ro",
|
||||
"sinhalese": "si",
|
||||
"castilian": "es",
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Tokenizer:
|
||||
"""A thin wrapper around `GPT2TokenizerFast` providing quick access to special tokens"""
|
||||
|
||||
tokenizer: "GPT2TokenizerFast"
|
||||
language: Optional[str]
|
||||
sot_sequence: Tuple[int]
|
||||
|
||||
def encode(self, text, **kwargs):
|
||||
return self.tokenizer.encode(text, **kwargs)
|
||||
|
||||
def decode(self, token_ids: Union[int, List[int], np.ndarray, torch.Tensor], **kwargs):
|
||||
return self.tokenizer.decode(token_ids, **kwargs)
|
||||
|
||||
def decode_with_timestamps(self, tokens) -> str:
|
||||
"""
|
||||
Timestamp tokens are above the special tokens' id range and are ignored by `decode()`.
|
||||
This method decodes given tokens with timestamps tokens annotated, e.g. "<|1.08|>".
|
||||
"""
|
||||
outputs = [[]]
|
||||
for token in tokens:
|
||||
if token >= self.timestamp_begin:
|
||||
timestamp = f"<|{(token - self.timestamp_begin) * 0.02:.2f}|>"
|
||||
outputs.append(timestamp)
|
||||
outputs.append([])
|
||||
else:
|
||||
outputs[-1].append(token)
|
||||
outputs = [s if isinstance(s, str) else self.tokenizer.decode(s) for s in outputs]
|
||||
return "".join(outputs)
|
||||
|
||||
@property
|
||||
@lru_cache()
|
||||
def eot(self) -> int:
|
||||
return self.tokenizer.eos_token_id
|
||||
|
||||
@property
|
||||
@lru_cache()
|
||||
def sot(self) -> int:
|
||||
return self._get_single_token_id("<|startoftranscript|>")
|
||||
|
||||
@property
|
||||
@lru_cache()
|
||||
def sot_lm(self) -> int:
|
||||
return self._get_single_token_id("<|startoflm|>")
|
||||
|
||||
@property
|
||||
@lru_cache()
|
||||
def sot_prev(self) -> int:
|
||||
return self._get_single_token_id("<|startofprev|>")
|
||||
|
||||
@property
|
||||
@lru_cache()
|
||||
def no_speech(self) -> int:
|
||||
return self._get_single_token_id("<|nospeech|>")
|
||||
|
||||
@property
|
||||
@lru_cache()
|
||||
def no_timestamps(self) -> int:
|
||||
return self._get_single_token_id("<|notimestamps|>")
|
||||
|
||||
@property
|
||||
@lru_cache()
|
||||
def timestamp_begin(self) -> int:
|
||||
return self.tokenizer.all_special_ids[-1] + 1
|
||||
|
||||
@property
|
||||
@lru_cache()
|
||||
def language_token(self) -> int:
|
||||
"""Returns the token id corresponding to the value of the `language` field"""
|
||||
if self.language is None:
|
||||
raise ValueError(f"This tokenizer does not have language token configured")
|
||||
|
||||
additional_tokens = dict(
|
||||
zip(
|
||||
self.tokenizer.additional_special_tokens,
|
||||
self.tokenizer.additional_special_tokens_ids,
|
||||
)
|
||||
)
|
||||
candidate = f"<|{self.language}|>"
|
||||
if candidate in additional_tokens:
|
||||
return additional_tokens[candidate]
|
||||
|
||||
raise KeyError(f"Language {self.language} not found in tokenizer.")
|
||||
|
||||
@property
|
||||
@lru_cache()
|
||||
def all_language_tokens(self) -> Tuple[int]:
|
||||
result = []
|
||||
for token, token_id in zip(
|
||||
self.tokenizer.additional_special_tokens,
|
||||
self.tokenizer.additional_special_tokens_ids,
|
||||
):
|
||||
if token.strip("<|>") in LANGUAGES:
|
||||
result.append(token_id)
|
||||
return tuple(result)
|
||||
|
||||
@property
|
||||
@lru_cache()
|
||||
def all_language_codes(self) -> Tuple[str]:
|
||||
return tuple(self.decode([l]).strip("<|>") for l in self.all_language_tokens)
|
||||
|
||||
@property
|
||||
@lru_cache()
|
||||
def sot_sequence_including_notimestamps(self) -> Tuple[int]:
|
||||
return tuple(list(self.sot_sequence) + [self.no_timestamps])
|
||||
|
||||
@property
|
||||
@lru_cache()
|
||||
def non_speech_tokens(self) -> Tuple[int]:
|
||||
"""
|
||||
Returns the list of tokens to suppress in order to avoid any speaker tags or non-speech
|
||||
annotations, to prevent sampling texts that are not actually spoken in the audio, e.g.
|
||||
|
||||
- ♪♪♪
|
||||
- ( SPEAKING FOREIGN LANGUAGE )
|
||||
- [DAVID] Hey there,
|
||||
|
||||
keeping basic punctuations like commas, periods, question marks, exclamation points, etc.
|
||||
"""
|
||||
symbols = list("\"#()*+/:;<=>@[\\]^_`{|}~「」『』")
|
||||
symbols += "<< >> <<< >>> -- --- -( -[ (' (\" (( )) ((( ))) [[ ]] {{ }} ♪♪ ♪♪♪".split()
|
||||
|
||||
# symbols that may be a single token or multiple tokens depending on the tokenizer.
|
||||
# In case they're multiple tokens, suppress the first token, which is safe because:
|
||||
# These are between U+2640 and U+267F miscellaneous symbols that are okay to suppress
|
||||
# in generations, and in the 3-byte UTF-8 representation they share the first two bytes.
|
||||
miscellaneous = set("♩♪♫♬♭♮♯")
|
||||
assert all(0x2640 <= ord(c) <= 0x267F for c in miscellaneous)
|
||||
|
||||
# allow hyphens "-" and single quotes "'" between words, but not at the beginning of a word
|
||||
result = {self.tokenizer.encode(" -")[0], self.tokenizer.encode(" '")[0]}
|
||||
for symbol in symbols + list(miscellaneous):
|
||||
for tokens in [self.tokenizer.encode(symbol), self.tokenizer.encode(" " + symbol)]:
|
||||
if len(tokens) == 1 or symbol in miscellaneous:
|
||||
result.add(tokens[0])
|
||||
|
||||
return tuple(sorted(result))
|
||||
|
||||
def _get_single_token_id(self, text) -> int:
|
||||
tokens = self.tokenizer.encode(text)
|
||||
assert len(tokens) == 1, f"{text} is not encoded as a single token"
|
||||
return tokens[0]
|
||||
|
||||
|
||||
@lru_cache(maxsize=None)
|
||||
def build_tokenizer(name: str = "gpt2"):
|
||||
os.environ["TOKENIZERS_PARALLELISM"] = "false"
|
||||
path = os.path.join(os.path.dirname(__file__), "assets", name)
|
||||
tokenizer = GPT2TokenizerFast.from_pretrained(path)
|
||||
|
||||
specials = [
|
||||
"<|startoftranscript|>",
|
||||
*[f"<|{lang}|>" for lang in LANGUAGES.keys()],
|
||||
"<|translate|>",
|
||||
"<|transcribe|>",
|
||||
"<|startoflm|>",
|
||||
"<|startofprev|>",
|
||||
"<|nospeech|>",
|
||||
"<|notimestamps|>",
|
||||
]
|
||||
|
||||
tokenizer.add_special_tokens(dict(additional_special_tokens=specials))
|
||||
return tokenizer
|
||||
|
||||
|
||||
@lru_cache(maxsize=None)
|
||||
def get_tokenizer(
|
||||
multilingual: bool,
|
||||
*,
|
||||
task: Optional[str] = None, # Literal["transcribe", "translate", None]
|
||||
language: Optional[str] = None,
|
||||
) -> Tokenizer:
|
||||
if language is not None:
|
||||
language = language.lower()
|
||||
if language not in LANGUAGES:
|
||||
if language in TO_LANGUAGE_CODE:
|
||||
language = TO_LANGUAGE_CODE[language]
|
||||
else:
|
||||
raise ValueError(f"Unsupported language: {language}")
|
||||
|
||||
if multilingual:
|
||||
tokenizer_name = "multilingual"
|
||||
task = task or "transcribe"
|
||||
language = language or "en"
|
||||
else:
|
||||
tokenizer_name = "gpt2"
|
||||
task = None
|
||||
language = None
|
||||
|
||||
tokenizer = build_tokenizer(name=tokenizer_name)
|
||||
all_special_ids: List[int] = tokenizer.all_special_ids
|
||||
sot: int = all_special_ids[1]
|
||||
translate: int = all_special_ids[-6]
|
||||
transcribe: int = all_special_ids[-5]
|
||||
|
||||
langs = tuple(LANGUAGES.keys())
|
||||
sot_sequence = [sot]
|
||||
if language is not None:
|
||||
sot_sequence.append(sot + 1 + langs.index(language))
|
||||
if task is not None:
|
||||
sot_sequence.append(transcribe if task == "transcribe" else translate)
|
||||
|
||||
return Tokenizer(tokenizer=tokenizer, language=language, sot_sequence=tuple(sot_sequence))
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from functools import lru_cache
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from transformers import GPT2TokenizerFast
|
||||
|
||||
LANGUAGES = {
|
||||
"en": "english",
|
||||
"zh": "chinese",
|
||||
"de": "german",
|
||||
"es": "spanish",
|
||||
"ru": "russian",
|
||||
"ko": "korean",
|
||||
"fr": "french",
|
||||
"ja": "japanese",
|
||||
"pt": "portuguese",
|
||||
"tr": "turkish",
|
||||
"pl": "polish",
|
||||
"ca": "catalan",
|
||||
"nl": "dutch",
|
||||
"ar": "arabic",
|
||||
"sv": "swedish",
|
||||
"it": "italian",
|
||||
"id": "indonesian",
|
||||
"hi": "hindi",
|
||||
"fi": "finnish",
|
||||
"vi": "vietnamese",
|
||||
"iw": "hebrew",
|
||||
"uk": "ukrainian",
|
||||
"el": "greek",
|
||||
"ms": "malay",
|
||||
"cs": "czech",
|
||||
"ro": "romanian",
|
||||
"da": "danish",
|
||||
"hu": "hungarian",
|
||||
"ta": "tamil",
|
||||
"no": "norwegian",
|
||||
"th": "thai",
|
||||
"ur": "urdu",
|
||||
"hr": "croatian",
|
||||
"bg": "bulgarian",
|
||||
"lt": "lithuanian",
|
||||
"la": "latin",
|
||||
"mi": "maori",
|
||||
"ml": "malayalam",
|
||||
"cy": "welsh",
|
||||
"sk": "slovak",
|
||||
"te": "telugu",
|
||||
"fa": "persian",
|
||||
"lv": "latvian",
|
||||
"bn": "bengali",
|
||||
"sr": "serbian",
|
||||
"az": "azerbaijani",
|
||||
"sl": "slovenian",
|
||||
"kn": "kannada",
|
||||
"et": "estonian",
|
||||
"mk": "macedonian",
|
||||
"br": "breton",
|
||||
"eu": "basque",
|
||||
"is": "icelandic",
|
||||
"hy": "armenian",
|
||||
"ne": "nepali",
|
||||
"mn": "mongolian",
|
||||
"bs": "bosnian",
|
||||
"kk": "kazakh",
|
||||
"sq": "albanian",
|
||||
"sw": "swahili",
|
||||
"gl": "galician",
|
||||
"mr": "marathi",
|
||||
"pa": "punjabi",
|
||||
"si": "sinhala",
|
||||
"km": "khmer",
|
||||
"sn": "shona",
|
||||
"yo": "yoruba",
|
||||
"so": "somali",
|
||||
"af": "afrikaans",
|
||||
"oc": "occitan",
|
||||
"ka": "georgian",
|
||||
"be": "belarusian",
|
||||
"tg": "tajik",
|
||||
"sd": "sindhi",
|
||||
"gu": "gujarati",
|
||||
"am": "amharic",
|
||||
"yi": "yiddish",
|
||||
"lo": "lao",
|
||||
"uz": "uzbek",
|
||||
"fo": "faroese",
|
||||
"ht": "haitian creole",
|
||||
"ps": "pashto",
|
||||
"tk": "turkmen",
|
||||
"nn": "nynorsk",
|
||||
"mt": "maltese",
|
||||
"sa": "sanskrit",
|
||||
"lb": "luxembourgish",
|
||||
"my": "myanmar",
|
||||
"bo": "tibetan",
|
||||
"tl": "tagalog",
|
||||
"mg": "malagasy",
|
||||
"as": "assamese",
|
||||
"tt": "tatar",
|
||||
"haw": "hawaiian",
|
||||
"ln": "lingala",
|
||||
"ha": "hausa",
|
||||
"ba": "bashkir",
|
||||
"jw": "javanese",
|
||||
"su": "sundanese",
|
||||
}
|
||||
|
||||
# language code lookup by name, with a few language aliases
|
||||
TO_LANGUAGE_CODE = {
|
||||
**{language: code for code, language in LANGUAGES.items()},
|
||||
"burmese": "my",
|
||||
"valencian": "ca",
|
||||
"flemish": "nl",
|
||||
"haitian": "ht",
|
||||
"letzeburgesch": "lb",
|
||||
"pushto": "ps",
|
||||
"panjabi": "pa",
|
||||
"moldavian": "ro",
|
||||
"moldovan": "ro",
|
||||
"sinhalese": "si",
|
||||
"castilian": "es",
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Tokenizer:
|
||||
"""A thin wrapper around `GPT2TokenizerFast` providing quick access to special tokens"""
|
||||
|
||||
tokenizer: "GPT2TokenizerFast"
|
||||
language: Optional[str]
|
||||
sot_sequence: Tuple[int]
|
||||
|
||||
def encode(self, text, **kwargs):
|
||||
return self.tokenizer.encode(text, **kwargs)
|
||||
|
||||
def decode(self, token_ids: Union[int, List[int], np.ndarray, torch.Tensor], **kwargs):
|
||||
return self.tokenizer.decode(token_ids, **kwargs)
|
||||
|
||||
def decode_with_timestamps(self, tokens) -> str:
|
||||
"""
|
||||
Timestamp tokens are above the special tokens' id range and are ignored by `decode()`.
|
||||
This method decodes given tokens with timestamps tokens annotated, e.g. "<|1.08|>".
|
||||
"""
|
||||
outputs = [[]]
|
||||
for token in tokens:
|
||||
if token >= self.timestamp_begin:
|
||||
timestamp = f"<|{(token - self.timestamp_begin) * 0.02:.2f}|>"
|
||||
outputs.append(timestamp)
|
||||
outputs.append([])
|
||||
else:
|
||||
outputs[-1].append(token)
|
||||
outputs = [s if isinstance(s, str) else self.tokenizer.decode(s) for s in outputs]
|
||||
return "".join(outputs)
|
||||
|
||||
@property
|
||||
@lru_cache()
|
||||
def eot(self) -> int:
|
||||
return self.tokenizer.eos_token_id
|
||||
|
||||
@property
|
||||
@lru_cache()
|
||||
def sot(self) -> int:
|
||||
return self._get_single_token_id("<|startoftranscript|>")
|
||||
|
||||
@property
|
||||
@lru_cache()
|
||||
def sot_lm(self) -> int:
|
||||
return self._get_single_token_id("<|startoflm|>")
|
||||
|
||||
@property
|
||||
@lru_cache()
|
||||
def sot_prev(self) -> int:
|
||||
return self._get_single_token_id("<|startofprev|>")
|
||||
|
||||
@property
|
||||
@lru_cache()
|
||||
def no_speech(self) -> int:
|
||||
return self._get_single_token_id("<|nospeech|>")
|
||||
|
||||
@property
|
||||
@lru_cache()
|
||||
def no_timestamps(self) -> int:
|
||||
return self._get_single_token_id("<|notimestamps|>")
|
||||
|
||||
@property
|
||||
@lru_cache()
|
||||
def timestamp_begin(self) -> int:
|
||||
return self.tokenizer.all_special_ids[-1] + 1
|
||||
|
||||
@property
|
||||
@lru_cache()
|
||||
def language_token(self) -> int:
|
||||
"""Returns the token id corresponding to the value of the `language` field"""
|
||||
if self.language is None:
|
||||
raise ValueError(f"This tokenizer does not have language token configured")
|
||||
|
||||
additional_tokens = dict(
|
||||
zip(
|
||||
self.tokenizer.additional_special_tokens,
|
||||
self.tokenizer.additional_special_tokens_ids,
|
||||
)
|
||||
)
|
||||
candidate = f"<|{self.language}|>"
|
||||
if candidate in additional_tokens:
|
||||
return additional_tokens[candidate]
|
||||
|
||||
raise KeyError(f"Language {self.language} not found in tokenizer.")
|
||||
|
||||
@property
|
||||
@lru_cache()
|
||||
def all_language_tokens(self) -> Tuple[int]:
|
||||
result = []
|
||||
for token, token_id in zip(
|
||||
self.tokenizer.additional_special_tokens,
|
||||
self.tokenizer.additional_special_tokens_ids,
|
||||
):
|
||||
if token.strip("<|>") in LANGUAGES:
|
||||
result.append(token_id)
|
||||
return tuple(result)
|
||||
|
||||
@property
|
||||
@lru_cache()
|
||||
def all_language_codes(self) -> Tuple[str]:
|
||||
return tuple(self.decode([l]).strip("<|>") for l in self.all_language_tokens)
|
||||
|
||||
@property
|
||||
@lru_cache()
|
||||
def sot_sequence_including_notimestamps(self) -> Tuple[int]:
|
||||
return tuple(list(self.sot_sequence) + [self.no_timestamps])
|
||||
|
||||
@property
|
||||
@lru_cache()
|
||||
def non_speech_tokens(self) -> Tuple[int]:
|
||||
"""
|
||||
Returns the list of tokens to suppress in order to avoid any speaker tags or non-speech
|
||||
annotations, to prevent sampling texts that are not actually spoken in the audio, e.g.
|
||||
|
||||
- ♪♪♪
|
||||
- ( SPEAKING FOREIGN LANGUAGE )
|
||||
- [DAVID] Hey there,
|
||||
|
||||
keeping basic punctuations like commas, periods, question marks, exclamation points, etc.
|
||||
"""
|
||||
symbols = list("\"#()*+/:;<=>@[\\]^_`{|}~「」『』")
|
||||
symbols += "<< >> <<< >>> -- --- -( -[ (' (\" (( )) ((( ))) [[ ]] {{ }} ♪♪ ♪♪♪".split()
|
||||
|
||||
# symbols that may be a single token or multiple tokens depending on the tokenizer.
|
||||
# In case they're multiple tokens, suppress the first token, which is safe because:
|
||||
# These are between U+2640 and U+267F miscellaneous symbols that are okay to suppress
|
||||
# in generations, and in the 3-byte UTF-8 representation they share the first two bytes.
|
||||
miscellaneous = set("♩♪♫♬♭♮♯")
|
||||
assert all(0x2640 <= ord(c) <= 0x267F for c in miscellaneous)
|
||||
|
||||
# allow hyphens "-" and single quotes "'" between words, but not at the beginning of a word
|
||||
result = {self.tokenizer.encode(" -")[0], self.tokenizer.encode(" '")[0]}
|
||||
for symbol in symbols + list(miscellaneous):
|
||||
for tokens in [self.tokenizer.encode(symbol), self.tokenizer.encode(" " + symbol)]:
|
||||
if len(tokens) == 1 or symbol in miscellaneous:
|
||||
result.add(tokens[0])
|
||||
|
||||
return tuple(sorted(result))
|
||||
|
||||
def _get_single_token_id(self, text) -> int:
|
||||
tokens = self.tokenizer.encode(text)
|
||||
assert len(tokens) == 1, f"{text} is not encoded as a single token"
|
||||
return tokens[0]
|
||||
|
||||
|
||||
@lru_cache(maxsize=None)
|
||||
def build_tokenizer(name: str = "gpt2"):
|
||||
os.environ["TOKENIZERS_PARALLELISM"] = "false"
|
||||
path = os.path.join(os.path.dirname(__file__), "assets", name)
|
||||
tokenizer = GPT2TokenizerFast.from_pretrained(path)
|
||||
|
||||
specials = [
|
||||
"<|startoftranscript|>",
|
||||
*[f"<|{lang}|>" for lang in LANGUAGES.keys()],
|
||||
"<|translate|>",
|
||||
"<|transcribe|>",
|
||||
"<|startoflm|>",
|
||||
"<|startofprev|>",
|
||||
"<|nospeech|>",
|
||||
"<|notimestamps|>",
|
||||
]
|
||||
|
||||
tokenizer.add_special_tokens(dict(additional_special_tokens=specials))
|
||||
return tokenizer
|
||||
|
||||
|
||||
@lru_cache(maxsize=None)
|
||||
def get_tokenizer(
|
||||
multilingual: bool,
|
||||
*,
|
||||
task: Optional[str] = None, # Literal["transcribe", "translate", None]
|
||||
language: Optional[str] = None,
|
||||
) -> Tokenizer:
|
||||
if language is not None:
|
||||
language = language.lower()
|
||||
if language not in LANGUAGES:
|
||||
if language in TO_LANGUAGE_CODE:
|
||||
language = TO_LANGUAGE_CODE[language]
|
||||
else:
|
||||
raise ValueError(f"Unsupported language: {language}")
|
||||
|
||||
if multilingual:
|
||||
tokenizer_name = "multilingual"
|
||||
task = task or "transcribe"
|
||||
language = language or "en"
|
||||
else:
|
||||
tokenizer_name = "gpt2"
|
||||
task = None
|
||||
language = None
|
||||
|
||||
tokenizer = build_tokenizer(name=tokenizer_name)
|
||||
all_special_ids: List[int] = tokenizer.all_special_ids
|
||||
sot: int = all_special_ids[1]
|
||||
translate: int = all_special_ids[-6]
|
||||
transcribe: int = all_special_ids[-5]
|
||||
|
||||
langs = tuple(LANGUAGES.keys())
|
||||
sot_sequence = [sot]
|
||||
if language is not None:
|
||||
sot_sequence.append(sot + 1 + langs.index(language))
|
||||
if task is not None:
|
||||
sot_sequence.append(transcribe if task == "transcribe" else translate)
|
||||
|
||||
return Tokenizer(tokenizer=tokenizer, language=language, sot_sequence=tuple(sot_sequence))
|
||||
|
||||
@@ -1,207 +1,207 @@
|
||||
import argparse
|
||||
import os
|
||||
import warnings
|
||||
from typing import List, Optional, Tuple, Union, TYPE_CHECKING
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import tqdm
|
||||
|
||||
from .audio import SAMPLE_RATE, N_FRAMES, HOP_LENGTH, pad_or_trim, log_mel_spectrogram
|
||||
from .decoding import DecodingOptions, DecodingResult
|
||||
from .tokenizer import LANGUAGES, TO_LANGUAGE_CODE, get_tokenizer
|
||||
from .utils import exact_div, format_timestamp, optional_int, optional_float, str2bool, write_txt, write_vtt, write_srt
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .model import Whisper
|
||||
|
||||
|
||||
def transcribe(
|
||||
model: "Whisper",
|
||||
audio: Union[str, np.ndarray, torch.Tensor],
|
||||
*,
|
||||
verbose: Optional[bool] = None,
|
||||
temperature: Union[float, Tuple[float, ...]] = (0.0, 0.2, 0.4, 0.6, 0.8, 1.0),
|
||||
compression_ratio_threshold: Optional[float] = 2.4,
|
||||
logprob_threshold: Optional[float] = -1.0,
|
||||
no_speech_threshold: Optional[float] = 0.6,
|
||||
condition_on_previous_text: bool = True,
|
||||
force_extraction: bool = False,
|
||||
**decode_options,
|
||||
):
|
||||
"""
|
||||
Transcribe an audio file using Whisper
|
||||
|
||||
Parameters
|
||||
----------
|
||||
model: Whisper
|
||||
The Whisper model instance
|
||||
|
||||
audio: Union[str, np.ndarray, torch.Tensor]
|
||||
The path to the audio file to open, or the audio waveform
|
||||
|
||||
verbose: bool
|
||||
Whether to display the text being decoded to the console. If True, displays all the details,
|
||||
If False, displays minimal details. If None, does not display anything
|
||||
|
||||
temperature: Union[float, Tuple[float, ...]]
|
||||
Temperature for sampling. It can be a tuple of temperatures, which will be successfully used
|
||||
upon failures according to either `compression_ratio_threshold` or `logprob_threshold`.
|
||||
|
||||
compression_ratio_threshold: float
|
||||
If the gzip compression ratio is above this value, treat as failed
|
||||
|
||||
logprob_threshold: float
|
||||
If the average log probability over sampled tokens is below this value, treat as failed
|
||||
|
||||
no_speech_threshold: float
|
||||
If the no_speech probability is higher than this value AND the average log probability
|
||||
over sampled tokens is below `logprob_threshold`, consider the segment as silent
|
||||
|
||||
condition_on_previous_text: bool
|
||||
if True, the previous output of the model is provided as a prompt for the next window;
|
||||
disabling may make the text inconsistent across windows, but the model becomes less prone to
|
||||
getting stuck in a failure loop, such as repetition looping or timestamps going out of sync.
|
||||
|
||||
decode_options: dict
|
||||
Keyword arguments to construct `DecodingOptions` instances
|
||||
|
||||
Returns
|
||||
-------
|
||||
A dictionary containing the resulting text ("text") and segment-level details ("segments"), and
|
||||
the spoken language ("language"), which is detected when `decode_options["language"]` is None.
|
||||
"""
|
||||
dtype = torch.float16 if decode_options.get("fp16", True) else torch.float32
|
||||
if model.device == torch.device("cpu"):
|
||||
if torch.cuda.is_available():
|
||||
warnings.warn("Performing inference on CPU when CUDA is available")
|
||||
if dtype == torch.float16:
|
||||
warnings.warn("FP16 is not supported on CPU; using FP32 instead")
|
||||
dtype = torch.float32
|
||||
|
||||
if dtype == torch.float32:
|
||||
decode_options["fp16"] = False
|
||||
|
||||
mel = log_mel_spectrogram(audio)
|
||||
|
||||
all_segments = []
|
||||
def add_segment(
|
||||
*, start: float, end: float, encoder_embeddings
|
||||
):
|
||||
|
||||
all_segments.append(
|
||||
{
|
||||
"start": start,
|
||||
"end": end,
|
||||
"encoder_embeddings":encoder_embeddings,
|
||||
}
|
||||
)
|
||||
# show the progress bar when verbose is False (otherwise the transcribed text will be printed)
|
||||
num_frames = mel.shape[-1]
|
||||
seek = 0
|
||||
previous_seek_value = seek
|
||||
sample_skip = 3000 #
|
||||
with tqdm.tqdm(total=num_frames, unit='frames', disable=verbose is not False) as pbar:
|
||||
while seek < num_frames:
|
||||
# seek是开始的帧数
|
||||
end_seek = min(seek + sample_skip, num_frames)
|
||||
segment = pad_or_trim(mel[:,seek:seek+sample_skip], N_FRAMES).to(model.device).to(dtype)
|
||||
|
||||
single = segment.ndim == 2
|
||||
if single:
|
||||
segment = segment.unsqueeze(0)
|
||||
if dtype == torch.float16:
|
||||
segment = segment.half()
|
||||
audio_features, embeddings = model.encoder(segment, include_embeddings = True)
|
||||
|
||||
encoder_embeddings = embeddings
|
||||
#print(f"encoder_embeddings shape {encoder_embeddings.shape}")
|
||||
add_segment(
|
||||
start=seek,
|
||||
end=end_seek,
|
||||
#text_tokens=tokens,
|
||||
#result=result,
|
||||
encoder_embeddings=encoder_embeddings,
|
||||
)
|
||||
seek+=sample_skip
|
||||
|
||||
return dict(segments=all_segments)
|
||||
|
||||
|
||||
def cli():
|
||||
from . import available_models
|
||||
|
||||
parser = argparse.ArgumentParser(formatter_class=argparse.ArgumentDefaultsHelpFormatter)
|
||||
parser.add_argument("audio", nargs="+", type=str, help="audio file(s) to transcribe")
|
||||
parser.add_argument("--model", default="small", choices=available_models(), help="name of the Whisper model to use")
|
||||
parser.add_argument("--model_dir", type=str, default=None, help="the path to save model files; uses ~/.cache/whisper by default")
|
||||
parser.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu", help="device to use for PyTorch inference")
|
||||
parser.add_argument("--output_dir", "-o", type=str, default=".", help="directory to save the outputs")
|
||||
parser.add_argument("--verbose", type=str2bool, default=True, help="whether to print out the progress and debug messages")
|
||||
|
||||
parser.add_argument("--task", type=str, default="transcribe", choices=["transcribe", "translate"], help="whether to perform X->X speech recognition ('transcribe') or X->English translation ('translate')")
|
||||
parser.add_argument("--language", type=str, default=None, choices=sorted(LANGUAGES.keys()) + sorted([k.title() for k in TO_LANGUAGE_CODE.keys()]), help="language spoken in the audio, specify None to perform language detection")
|
||||
|
||||
parser.add_argument("--temperature", type=float, default=0, help="temperature to use for sampling")
|
||||
parser.add_argument("--best_of", type=optional_int, default=5, help="number of candidates when sampling with non-zero temperature")
|
||||
parser.add_argument("--beam_size", type=optional_int, default=5, help="number of beams in beam search, only applicable when temperature is zero")
|
||||
parser.add_argument("--patience", type=float, default=None, help="optional patience value to use in beam decoding, as in https://arxiv.org/abs/2204.05424, the default (1.0) is equivalent to conventional beam search")
|
||||
parser.add_argument("--length_penalty", type=float, default=None, help="optional token length penalty coefficient (alpha) as in https://arxiv.org/abs/1609.08144, uses simple length normalization by default")
|
||||
|
||||
parser.add_argument("--suppress_tokens", type=str, default="-1", help="comma-separated list of token ids to suppress during sampling; '-1' will suppress most special characters except common punctuations")
|
||||
parser.add_argument("--initial_prompt", type=str, default=None, help="optional text to provide as a prompt for the first window.")
|
||||
parser.add_argument("--condition_on_previous_text", type=str2bool, default=True, help="if True, provide the previous output of the model as a prompt for the next window; disabling may make the text inconsistent across windows, but the model becomes less prone to getting stuck in a failure loop")
|
||||
parser.add_argument("--fp16", type=str2bool, default=True, help="whether to perform inference in fp16; True by default")
|
||||
|
||||
parser.add_argument("--temperature_increment_on_fallback", type=optional_float, default=0.2, help="temperature to increase when falling back when the decoding fails to meet either of the thresholds below")
|
||||
parser.add_argument("--compression_ratio_threshold", type=optional_float, default=2.4, help="if the gzip compression ratio is higher than this value, treat the decoding as failed")
|
||||
parser.add_argument("--logprob_threshold", type=optional_float, default=-1.0, help="if the average log probability is lower than this value, treat the decoding as failed")
|
||||
parser.add_argument("--no_speech_threshold", type=optional_float, default=0.6, help="if the probability of the <|nospeech|> token is higher than this value AND the decoding has failed due to `logprob_threshold`, consider the segment as silence")
|
||||
parser.add_argument("--threads", type=optional_int, default=0, help="number of threads used by torch for CPU inference; supercedes MKL_NUM_THREADS/OMP_NUM_THREADS")
|
||||
|
||||
args = parser.parse_args().__dict__
|
||||
model_name: str = args.pop("model")
|
||||
model_dir: str = args.pop("model_dir")
|
||||
output_dir: str = args.pop("output_dir")
|
||||
device: str = args.pop("device")
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
if model_name.endswith(".en") and args["language"] not in {"en", "English"}:
|
||||
if args["language"] is not None:
|
||||
warnings.warn(f"{model_name} is an English-only model but receipted '{args['language']}'; using English instead.")
|
||||
args["language"] = "en"
|
||||
|
||||
temperature = args.pop("temperature")
|
||||
temperature_increment_on_fallback = args.pop("temperature_increment_on_fallback")
|
||||
if temperature_increment_on_fallback is not None:
|
||||
temperature = tuple(np.arange(temperature, 1.0 + 1e-6, temperature_increment_on_fallback))
|
||||
else:
|
||||
temperature = [temperature]
|
||||
|
||||
threads = args.pop("threads")
|
||||
if threads > 0:
|
||||
torch.set_num_threads(threads)
|
||||
|
||||
from . import load_model
|
||||
model = load_model(model_name, device=device, download_root=model_dir)
|
||||
|
||||
for audio_path in args.pop("audio"):
|
||||
result = transcribe(model, audio_path, temperature=temperature, **args)
|
||||
|
||||
audio_basename = os.path.basename(audio_path)
|
||||
|
||||
# save TXT
|
||||
with open(os.path.join(output_dir, audio_basename + ".txt"), "w", encoding="utf-8") as txt:
|
||||
write_txt(result["segments"], file=txt)
|
||||
|
||||
# save VTT
|
||||
with open(os.path.join(output_dir, audio_basename + ".vtt"), "w", encoding="utf-8") as vtt:
|
||||
write_vtt(result["segments"], file=vtt)
|
||||
|
||||
# save SRT
|
||||
with open(os.path.join(output_dir, audio_basename + ".srt"), "w", encoding="utf-8") as srt:
|
||||
write_srt(result["segments"], file=srt)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
cli()
|
||||
import argparse
|
||||
import os
|
||||
import warnings
|
||||
from typing import List, Optional, Tuple, Union, TYPE_CHECKING
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import tqdm
|
||||
|
||||
from .audio import SAMPLE_RATE, N_FRAMES, HOP_LENGTH, pad_or_trim, log_mel_spectrogram
|
||||
from .decoding import DecodingOptions, DecodingResult
|
||||
from .tokenizer import LANGUAGES, TO_LANGUAGE_CODE, get_tokenizer
|
||||
from .utils import exact_div, format_timestamp, optional_int, optional_float, str2bool, write_txt, write_vtt, write_srt
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .model import Whisper
|
||||
|
||||
|
||||
def transcribe(
|
||||
model: "Whisper",
|
||||
audio: Union[str, np.ndarray, torch.Tensor],
|
||||
*,
|
||||
verbose: Optional[bool] = None,
|
||||
temperature: Union[float, Tuple[float, ...]] = (0.0, 0.2, 0.4, 0.6, 0.8, 1.0),
|
||||
compression_ratio_threshold: Optional[float] = 2.4,
|
||||
logprob_threshold: Optional[float] = -1.0,
|
||||
no_speech_threshold: Optional[float] = 0.6,
|
||||
condition_on_previous_text: bool = True,
|
||||
force_extraction: bool = False,
|
||||
**decode_options,
|
||||
):
|
||||
"""
|
||||
Transcribe an audio file using Whisper
|
||||
|
||||
Parameters
|
||||
----------
|
||||
model: Whisper
|
||||
The Whisper model instance
|
||||
|
||||
audio: Union[str, np.ndarray, torch.Tensor]
|
||||
The path to the audio file to open, or the audio waveform
|
||||
|
||||
verbose: bool
|
||||
Whether to display the text being decoded to the console. If True, displays all the details,
|
||||
If False, displays minimal details. If None, does not display anything
|
||||
|
||||
temperature: Union[float, Tuple[float, ...]]
|
||||
Temperature for sampling. It can be a tuple of temperatures, which will be successfully used
|
||||
upon failures according to either `compression_ratio_threshold` or `logprob_threshold`.
|
||||
|
||||
compression_ratio_threshold: float
|
||||
If the gzip compression ratio is above this value, treat as failed
|
||||
|
||||
logprob_threshold: float
|
||||
If the average log probability over sampled tokens is below this value, treat as failed
|
||||
|
||||
no_speech_threshold: float
|
||||
If the no_speech probability is higher than this value AND the average log probability
|
||||
over sampled tokens is below `logprob_threshold`, consider the segment as silent
|
||||
|
||||
condition_on_previous_text: bool
|
||||
if True, the previous output of the model is provided as a prompt for the next window;
|
||||
disabling may make the text inconsistent across windows, but the model becomes less prone to
|
||||
getting stuck in a failure loop, such as repetition looping or timestamps going out of sync.
|
||||
|
||||
decode_options: dict
|
||||
Keyword arguments to construct `DecodingOptions` instances
|
||||
|
||||
Returns
|
||||
-------
|
||||
A dictionary containing the resulting text ("text") and segment-level details ("segments"), and
|
||||
the spoken language ("language"), which is detected when `decode_options["language"]` is None.
|
||||
"""
|
||||
dtype = torch.float16 if decode_options.get("fp16", True) else torch.float32
|
||||
if model.device == torch.device("cpu"):
|
||||
if torch.cuda.is_available():
|
||||
warnings.warn("Performing inference on CPU when CUDA is available")
|
||||
if dtype == torch.float16:
|
||||
warnings.warn("FP16 is not supported on CPU; using FP32 instead")
|
||||
dtype = torch.float32
|
||||
|
||||
if dtype == torch.float32:
|
||||
decode_options["fp16"] = False
|
||||
|
||||
mel = log_mel_spectrogram(audio)
|
||||
|
||||
all_segments = []
|
||||
def add_segment(
|
||||
*, start: float, end: float, encoder_embeddings
|
||||
):
|
||||
|
||||
all_segments.append(
|
||||
{
|
||||
"start": start,
|
||||
"end": end,
|
||||
"encoder_embeddings":encoder_embeddings,
|
||||
}
|
||||
)
|
||||
# show the progress bar when verbose is False (otherwise the transcribed text will be printed)
|
||||
num_frames = mel.shape[-1]
|
||||
seek = 0
|
||||
previous_seek_value = seek
|
||||
sample_skip = 3000 #
|
||||
with tqdm.tqdm(total=num_frames, unit='frames', disable=verbose is not False) as pbar:
|
||||
while seek < num_frames:
|
||||
# seek是开始的帧数
|
||||
end_seek = min(seek + sample_skip, num_frames)
|
||||
segment = pad_or_trim(mel[:,seek:seek+sample_skip], N_FRAMES).to(model.device).to(dtype)
|
||||
|
||||
single = segment.ndim == 2
|
||||
if single:
|
||||
segment = segment.unsqueeze(0)
|
||||
if dtype == torch.float16:
|
||||
segment = segment.half()
|
||||
audio_features, embeddings = model.encoder(segment, include_embeddings = True)
|
||||
|
||||
encoder_embeddings = embeddings
|
||||
#print(f"encoder_embeddings shape {encoder_embeddings.shape}")
|
||||
add_segment(
|
||||
start=seek,
|
||||
end=end_seek,
|
||||
#text_tokens=tokens,
|
||||
#result=result,
|
||||
encoder_embeddings=encoder_embeddings,
|
||||
)
|
||||
seek+=sample_skip
|
||||
|
||||
return dict(segments=all_segments)
|
||||
|
||||
|
||||
def cli():
|
||||
from . import available_models
|
||||
|
||||
parser = argparse.ArgumentParser(formatter_class=argparse.ArgumentDefaultsHelpFormatter)
|
||||
parser.add_argument("audio", nargs="+", type=str, help="audio file(s) to transcribe")
|
||||
parser.add_argument("--model", default="small", choices=available_models(), help="name of the Whisper model to use")
|
||||
parser.add_argument("--model_dir", type=str, default=None, help="the path to save model files; uses ~/.cache/whisper by default")
|
||||
parser.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu", help="device to use for PyTorch inference")
|
||||
parser.add_argument("--output_dir", "-o", type=str, default=".", help="directory to save the outputs")
|
||||
parser.add_argument("--verbose", type=str2bool, default=True, help="whether to print out the progress and debug messages")
|
||||
|
||||
parser.add_argument("--task", type=str, default="transcribe", choices=["transcribe", "translate"], help="whether to perform X->X speech recognition ('transcribe') or X->English translation ('translate')")
|
||||
parser.add_argument("--language", type=str, default=None, choices=sorted(LANGUAGES.keys()) + sorted([k.title() for k in TO_LANGUAGE_CODE.keys()]), help="language spoken in the audio, specify None to perform language detection")
|
||||
|
||||
parser.add_argument("--temperature", type=float, default=0, help="temperature to use for sampling")
|
||||
parser.add_argument("--best_of", type=optional_int, default=5, help="number of candidates when sampling with non-zero temperature")
|
||||
parser.add_argument("--beam_size", type=optional_int, default=5, help="number of beams in beam search, only applicable when temperature is zero")
|
||||
parser.add_argument("--patience", type=float, default=None, help="optional patience value to use in beam decoding, as in https://arxiv.org/abs/2204.05424, the default (1.0) is equivalent to conventional beam search")
|
||||
parser.add_argument("--length_penalty", type=float, default=None, help="optional token length penalty coefficient (alpha) as in https://arxiv.org/abs/1609.08144, uses simple length normalization by default")
|
||||
|
||||
parser.add_argument("--suppress_tokens", type=str, default="-1", help="comma-separated list of token ids to suppress during sampling; '-1' will suppress most special characters except common punctuations")
|
||||
parser.add_argument("--initial_prompt", type=str, default=None, help="optional text to provide as a prompt for the first window.")
|
||||
parser.add_argument("--condition_on_previous_text", type=str2bool, default=True, help="if True, provide the previous output of the model as a prompt for the next window; disabling may make the text inconsistent across windows, but the model becomes less prone to getting stuck in a failure loop")
|
||||
parser.add_argument("--fp16", type=str2bool, default=True, help="whether to perform inference in fp16; True by default")
|
||||
|
||||
parser.add_argument("--temperature_increment_on_fallback", type=optional_float, default=0.2, help="temperature to increase when falling back when the decoding fails to meet either of the thresholds below")
|
||||
parser.add_argument("--compression_ratio_threshold", type=optional_float, default=2.4, help="if the gzip compression ratio is higher than this value, treat the decoding as failed")
|
||||
parser.add_argument("--logprob_threshold", type=optional_float, default=-1.0, help="if the average log probability is lower than this value, treat the decoding as failed")
|
||||
parser.add_argument("--no_speech_threshold", type=optional_float, default=0.6, help="if the probability of the <|nospeech|> token is higher than this value AND the decoding has failed due to `logprob_threshold`, consider the segment as silence")
|
||||
parser.add_argument("--threads", type=optional_int, default=0, help="number of threads used by torch for CPU inference; supercedes MKL_NUM_THREADS/OMP_NUM_THREADS")
|
||||
|
||||
args = parser.parse_args().__dict__
|
||||
model_name: str = args.pop("model")
|
||||
model_dir: str = args.pop("model_dir")
|
||||
output_dir: str = args.pop("output_dir")
|
||||
device: str = args.pop("device")
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
if model_name.endswith(".en") and args["language"] not in {"en", "English"}:
|
||||
if args["language"] is not None:
|
||||
warnings.warn(f"{model_name} is an English-only model but receipted '{args['language']}'; using English instead.")
|
||||
args["language"] = "en"
|
||||
|
||||
temperature = args.pop("temperature")
|
||||
temperature_increment_on_fallback = args.pop("temperature_increment_on_fallback")
|
||||
if temperature_increment_on_fallback is not None:
|
||||
temperature = tuple(np.arange(temperature, 1.0 + 1e-6, temperature_increment_on_fallback))
|
||||
else:
|
||||
temperature = [temperature]
|
||||
|
||||
threads = args.pop("threads")
|
||||
if threads > 0:
|
||||
torch.set_num_threads(threads)
|
||||
|
||||
from . import load_model
|
||||
model = load_model(model_name, device=device, download_root=model_dir)
|
||||
|
||||
for audio_path in args.pop("audio"):
|
||||
result = transcribe(model, audio_path, temperature=temperature, **args)
|
||||
|
||||
audio_basename = os.path.basename(audio_path)
|
||||
|
||||
# save TXT
|
||||
with open(os.path.join(output_dir, audio_basename + ".txt"), "w", encoding="utf-8") as txt:
|
||||
write_txt(result["segments"], file=txt)
|
||||
|
||||
# save VTT
|
||||
with open(os.path.join(output_dir, audio_basename + ".vtt"), "w", encoding="utf-8") as vtt:
|
||||
write_vtt(result["segments"], file=vtt)
|
||||
|
||||
# save SRT
|
||||
with open(os.path.join(output_dir, audio_basename + ".srt"), "w", encoding="utf-8") as srt:
|
||||
write_srt(result["segments"], file=srt)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
cli()
|
||||
|
||||
@@ -1,87 +1,87 @@
|
||||
import zlib
|
||||
from typing import Iterator, TextIO
|
||||
|
||||
|
||||
def exact_div(x, y):
|
||||
assert x % y == 0
|
||||
return x // y
|
||||
|
||||
|
||||
def str2bool(string):
|
||||
str2val = {"True": True, "False": False}
|
||||
if string in str2val:
|
||||
return str2val[string]
|
||||
else:
|
||||
raise ValueError(f"Expected one of {set(str2val.keys())}, got {string}")
|
||||
|
||||
|
||||
def optional_int(string):
|
||||
return None if string == "None" else int(string)
|
||||
|
||||
|
||||
def optional_float(string):
|
||||
return None if string == "None" else float(string)
|
||||
|
||||
|
||||
def compression_ratio(text) -> float:
|
||||
return len(text) / len(zlib.compress(text.encode("utf-8")))
|
||||
|
||||
|
||||
def format_timestamp(seconds: float, always_include_hours: bool = False, decimal_marker: str = '.'):
|
||||
assert seconds >= 0, "non-negative timestamp expected"
|
||||
milliseconds = round(seconds * 1000.0)
|
||||
|
||||
hours = milliseconds // 3_600_000
|
||||
milliseconds -= hours * 3_600_000
|
||||
|
||||
minutes = milliseconds // 60_000
|
||||
milliseconds -= minutes * 60_000
|
||||
|
||||
seconds = milliseconds // 1_000
|
||||
milliseconds -= seconds * 1_000
|
||||
|
||||
hours_marker = f"{hours:02d}:" if always_include_hours or hours > 0 else ""
|
||||
return f"{hours_marker}{minutes:02d}:{seconds:02d}{decimal_marker}{milliseconds:03d}"
|
||||
|
||||
|
||||
def write_txt(transcript: Iterator[dict], file: TextIO):
|
||||
for segment in transcript:
|
||||
print(segment['text'].strip(), file=file, flush=True)
|
||||
|
||||
|
||||
def write_vtt(transcript: Iterator[dict], file: TextIO):
|
||||
print("WEBVTT\n", file=file)
|
||||
for segment in transcript:
|
||||
print(
|
||||
f"{format_timestamp(segment['start'])} --> {format_timestamp(segment['end'])}\n"
|
||||
f"{segment['text'].strip().replace('-->', '->')}\n",
|
||||
file=file,
|
||||
flush=True,
|
||||
)
|
||||
|
||||
|
||||
def write_srt(transcript: Iterator[dict], file: TextIO):
|
||||
"""
|
||||
Write a transcript to a file in SRT format.
|
||||
|
||||
Example usage:
|
||||
from pathlib import Path
|
||||
from whisper.utils import write_srt
|
||||
|
||||
result = transcribe(model, audio_path, temperature=temperature, **args)
|
||||
|
||||
# save SRT
|
||||
audio_basename = Path(audio_path).stem
|
||||
with open(Path(output_dir) / (audio_basename + ".srt"), "w", encoding="utf-8") as srt:
|
||||
write_srt(result["segments"], file=srt)
|
||||
"""
|
||||
for i, segment in enumerate(transcript, start=1):
|
||||
# write srt lines
|
||||
print(
|
||||
f"{i}\n"
|
||||
f"{format_timestamp(segment['start'], always_include_hours=True, decimal_marker=',')} --> "
|
||||
f"{format_timestamp(segment['end'], always_include_hours=True, decimal_marker=',')}\n"
|
||||
f"{segment['text'].strip().replace('-->', '->')}\n",
|
||||
file=file,
|
||||
flush=True,
|
||||
)
|
||||
import zlib
|
||||
from typing import Iterator, TextIO
|
||||
|
||||
|
||||
def exact_div(x, y):
|
||||
assert x % y == 0
|
||||
return x // y
|
||||
|
||||
|
||||
def str2bool(string):
|
||||
str2val = {"True": True, "False": False}
|
||||
if string in str2val:
|
||||
return str2val[string]
|
||||
else:
|
||||
raise ValueError(f"Expected one of {set(str2val.keys())}, got {string}")
|
||||
|
||||
|
||||
def optional_int(string):
|
||||
return None if string == "None" else int(string)
|
||||
|
||||
|
||||
def optional_float(string):
|
||||
return None if string == "None" else float(string)
|
||||
|
||||
|
||||
def compression_ratio(text) -> float:
|
||||
return len(text) / len(zlib.compress(text.encode("utf-8")))
|
||||
|
||||
|
||||
def format_timestamp(seconds: float, always_include_hours: bool = False, decimal_marker: str = '.'):
|
||||
assert seconds >= 0, "non-negative timestamp expected"
|
||||
milliseconds = round(seconds * 1000.0)
|
||||
|
||||
hours = milliseconds // 3_600_000
|
||||
milliseconds -= hours * 3_600_000
|
||||
|
||||
minutes = milliseconds // 60_000
|
||||
milliseconds -= minutes * 60_000
|
||||
|
||||
seconds = milliseconds // 1_000
|
||||
milliseconds -= seconds * 1_000
|
||||
|
||||
hours_marker = f"{hours:02d}:" if always_include_hours or hours > 0 else ""
|
||||
return f"{hours_marker}{minutes:02d}:{seconds:02d}{decimal_marker}{milliseconds:03d}"
|
||||
|
||||
|
||||
def write_txt(transcript: Iterator[dict], file: TextIO):
|
||||
for segment in transcript:
|
||||
print(segment['text'].strip(), file=file, flush=True)
|
||||
|
||||
|
||||
def write_vtt(transcript: Iterator[dict], file: TextIO):
|
||||
print("WEBVTT\n", file=file)
|
||||
for segment in transcript:
|
||||
print(
|
||||
f"{format_timestamp(segment['start'])} --> {format_timestamp(segment['end'])}\n"
|
||||
f"{segment['text'].strip().replace('-->', '->')}\n",
|
||||
file=file,
|
||||
flush=True,
|
||||
)
|
||||
|
||||
|
||||
def write_srt(transcript: Iterator[dict], file: TextIO):
|
||||
"""
|
||||
Write a transcript to a file in SRT format.
|
||||
|
||||
Example usage:
|
||||
from pathlib import Path
|
||||
from whisper.utils import write_srt
|
||||
|
||||
result = transcribe(model, audio_path, temperature=temperature, **args)
|
||||
|
||||
# save SRT
|
||||
audio_basename = Path(audio_path).stem
|
||||
with open(Path(output_dir) / (audio_basename + ".srt"), "w", encoding="utf-8") as srt:
|
||||
write_srt(result["segments"], file=srt)
|
||||
"""
|
||||
for i, segment in enumerate(transcript, start=1):
|
||||
# write srt lines
|
||||
print(
|
||||
f"{i}\n"
|
||||
f"{format_timestamp(segment['start'], always_include_hours=True, decimal_marker=',')} --> "
|
||||
f"{format_timestamp(segment['end'], always_include_hours=True, decimal_marker=',')}\n"
|
||||
f"{segment['text'].strip().replace('-->', '->')}\n",
|
||||
file=file,
|
||||
flush=True,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user