Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
109 changes: 56 additions & 53 deletions src/diffusers/utils/loading_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -85,61 +85,64 @@ def load_video(
f"Incorrect path or URL. URLs must start with `http://` or `https://`, and {video} is not a valid path."
)

if is_url:
response = requests.get(video, stream=True)
if response.status_code != 200:
raise ValueError(f"Failed to download video. Status code: {response.status_code}")

parsed_url = urlparse(video)
file_name = os.path.basename(unquote(parsed_url.path))

suffix = os.path.splitext(file_name)[1] or ".mp4"
video_path = tempfile.NamedTemporaryFile(suffix=suffix, delete=False).name

was_tempfile_created = True

video_data = response.iter_content(chunk_size=8192)
with open(video_path, "wb") as f:
for chunk in video_data:
f.write(chunk)

video = video_path

pil_images = []
fps = None
if video.endswith(".gif"):
gif = PIL.Image.open(video)
# Milliseconds this frame is displayed for; GIFs are not obliged to record it.
frame_duration = gif.info.get("duration")
fps = 1000 / frame_duration if frame_duration else None
try:
while True:
pil_images.append(gif.copy())
gif.seek(gif.tell() + 1)
except EOFError:
pass
try:
if is_url:
response = requests.get(video, stream=True)
if response.status_code != 200:
raise ValueError(f"Failed to download video. Status code: {response.status_code}")

parsed_url = urlparse(video)
file_name = os.path.basename(unquote(parsed_url.path))

suffix = os.path.splitext(file_name)[1] or ".mp4"
with tempfile.NamedTemporaryFile(suffix=suffix, delete=False) as tmp:
video_path = tmp.name

was_tempfile_created = True

video_data = response.iter_content(chunk_size=8192)
with open(video_path, "wb") as f:
for chunk in video_data:
f.write(chunk)

video = video_path

pil_images = []
fps = None
if video.endswith(".gif"):
with PIL.Image.open(video) as gif:
# Milliseconds this frame is displayed for; GIFs are not obliged to record it.
frame_duration = gif.info.get("duration")
fps = 1000 / frame_duration if frame_duration else None
try:
while True:
pil_images.append(gif.copy())
gif.seek(gif.tell() + 1)
except EOFError:
pass

else:
if is_imageio_available():
import imageio
else:
raise ImportError(BACKENDS_MAPPING["imageio"][1].format("load_video"))

try:
imageio.plugins.ffmpeg.get_exe()
except AttributeError:
raise AttributeError(
"`Unable to find an ffmpeg installation on your machine. Please install via `pip install imageio-ffmpeg"
)

with imageio.get_reader(video) as reader:
fps = reader.get_meta_data().get("fps")
# Read all frames
for frame in reader:
pil_images.append(PIL.Image.fromarray(frame))

if was_tempfile_created:
os.remove(video_path)
if is_imageio_available():
import imageio
else:
raise ImportError(BACKENDS_MAPPING["imageio"][1].format("load_video"))

try:
imageio.plugins.ffmpeg.get_exe()
except AttributeError:
raise AttributeError(
"`Unable to find an ffmpeg installation on your machine. Please install via `pip install imageio-ffmpeg"
)

with imageio.get_reader(video) as reader:
fps = reader.get_meta_data().get("fps")
# Read all frames
for frame in reader:
pil_images.append(PIL.Image.fromarray(frame))

finally:
if was_tempfile_created:
os.remove(video_path)

if convert_method is not None:
pil_images = convert_method(pil_images)
Expand Down
70 changes: 70 additions & 0 deletions tests/others/test_loading_utils.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,70 @@
from io import BytesIO

import pytest
from PIL import Image, UnidentifiedImageError

from diffusers.utils import load_video


class FakeResponse:
status_code = 200

def __init__(self, data, interrupt=False):
self.data = data
self.interrupt = interrupt

def iter_content(self, chunk_size=8192):
yield self.data
if self.interrupt:
raise ConnectionError("Download interrupted")


def mock_remote(monkeypatch, tmp_path, response):
monkeypatch.setattr(
"diffusers.utils.loading_utils.requests.get",
lambda *args, **kwargs: response,
)
monkeypatch.setattr(
"diffusers.utils.loading_utils.tempfile.tempdir",
str(tmp_path),
)


@pytest.mark.parametrize("interrupt", [False, True])
def test_remote_video_failure_cleans_tempfile(tmp_path, monkeypatch, interrupt):
response = FakeResponse(b"invalid-gif-data", interrupt=interrupt)
mock_remote(monkeypatch, tmp_path, response)

expected = ConnectionError if interrupt else UnidentifiedImageError

with pytest.raises(expected):
load_video("https://example.com/broken.gif")

assert not list(tmp_path.iterdir())


def test_remote_video_success_cleans_tempfile(tmp_path, monkeypatch):
buffer = BytesIO()
Image.new("RGB", (4, 4), "red").save(
buffer, format="GIF", duration=100
)

mock_remote(monkeypatch, tmp_path, FakeResponse(buffer.getvalue()))

frames, fps = load_video(
"https://example.com/valid.gif", return_fps=True
)

assert len(frames) == 1
assert fps == pytest.approx(10.0)
assert not list(tmp_path.iterdir())


def test_local_video_is_preserved(tmp_path):
video_path = tmp_path / "local.gif"
Image.new("RGB", (4, 4), "red").save(video_path, format="GIF")

frames = load_video(str(video_path))

assert len(frames) == 1
assert video_path.exists()
Loading