Refactor new types out of the base comfy folder

I've created a new `comfy_api` top-level module. The intention is that
anything within this folder would be covered by semver-style versioning
that would allow custom nodes to rely on them not introducing breaking
changes.
This commit is contained in:
Jacob Segal 2025-04-28 19:30:30 -07:00
parent 5c55ce5938
commit 08f5d7d0b7
9 changed files with 172 additions and 140 deletions

View File

@ -1,7 +1,6 @@
import torch import torch
from typing import Callable, Protocol, TypedDict, Optional, List from typing import Callable, Protocol, TypedDict, Optional, List
from .node_typing import IO, InputTypeDict, ComfyNodeABC, CheckLazyMixin, FileLocator from .node_typing import IO, InputTypeDict, ComfyNodeABC, CheckLazyMixin, FileLocator
from .input_types import ImageInput, AudioInput, VideoInput, VideoFromFile, VideoFromComponents, VideoContainer, VideoCodec, VideoComponents
class UnetApplyFunction(Protocol): class UnetApplyFunction(Protocol):
@ -44,12 +43,4 @@ __all__ = [
ComfyNodeABC.__name__, ComfyNodeABC.__name__,
CheckLazyMixin.__name__, CheckLazyMixin.__name__,
FileLocator.__name__, FileLocator.__name__,
ImageInput.__name__,
AudioInput.__name__,
VideoInput.__name__,
VideoFromFile.__name__,
VideoFromComponents.__name__,
VideoContainer.__name__,
VideoCodec.__name__,
VideoComponents.__name__,
] ]

View File

@ -0,0 +1,8 @@
from .basic_types import ImageInput, AudioInput
from .video_types import VideoInput
__all__ = [
"ImageInput",
"AudioInput",
"VideoInput",
]

View File

@ -0,0 +1,20 @@
import torch
from typing import TypedDict
ImageInput = torch.Tensor
"""
An image in format [B, H, W, C] where B is the batch size, C is the number of channels,
"""
class AudioInput(TypedDict):
"""
TypedDict representing audio input.
"""
waveform: torch.Tensor
"""
Tensor in the format [B, C, T] where B is the batch size, C is the number of channels,
"""
sample_rate: int

View File

@ -0,0 +1,45 @@
from __future__ import annotations
from abc import ABC, abstractmethod
from typing import Optional
from comfy_api.util import VideoContainer, VideoCodec
class VideoInput(ABC):
"""
Abstract base class for video input types.
"""
@abstractmethod
def get_components(self) -> VideoComponents:
"""
Abstract method to get the video components (images, audio, and frame rate).
Returns:
VideoComponents containing images, audio, and frame rate
"""
pass
@abstractmethod
def save_to(
self,
path: str,
format: VideoContainer = VideoContainer.AUTO,
codec: VideoCodec = VideoCodec.AUTO,
metadata: Optional[dict] = None
):
"""
Abstract method to save the video input to a file.
"""
pass
# Provide a default implementation, but subclasses can provide optimized versions
# if possible.
def get_dimensions(self) -> tuple[int, int]:
"""
Returns the dimensions of the video input.
Returns:
Tuple of (width, height)
"""
components = self.get_components()
return components.images.shape[2], components.images.shape[1]

View File

@ -0,0 +1,7 @@
from .video_types import VideoFromFile, VideoFromComponents
__all__ = [
# Implementations
"VideoFromFile",
"VideoFromComponents",
]

View File

@ -1,117 +1,16 @@
from __future__ import annotations from __future__ import annotations
from abc import ABC, abstractmethod
from av.container import InputContainer from av.container import InputContainer
from av.subtitles.stream import SubtitleStream from av.subtitles.stream import SubtitleStream
from dataclasses import dataclass
from enum import Enum
from fractions import Fraction from fractions import Fraction
from typing import Optional, TypedDict from typing import Optional
from comfy_api.input import AudioInput
import av import av
import io import io
import json import json
import numpy as np import numpy as np
import torch import torch
from comfy_api.input import VideoInput
ImageInput = torch.Tensor from comfy_api.util import VideoContainer, VideoCodec, VideoComponents
"""
An image in format [B, H, W, C] where B is the batch size, C is the number of channels,
"""
class AudioInput(TypedDict):
"""
TypedDict representing audio input.
"""
waveform: torch.Tensor
"""
Tensor in the format [B, C, T] where B is the batch size, C is the number of channels,
"""
sample_rate: int
class VideoCodec(str, Enum):
AUTO = "auto"
H264 = "h264"
@classmethod
def as_input(cls) -> list[str]:
"""
Returns a list of codec names that can be used as node input.
"""
return [member.value for member in cls]
class VideoContainer(str, Enum):
AUTO = "auto"
MP4 = "mp4"
@classmethod
def as_input(cls) -> list[str]:
"""
Returns a list of container names that can be used as node input.
"""
return [member.value for member in cls]
@classmethod
def get_extension(cls, value) -> str:
"""
Returns the file extension for the container.
"""
if isinstance(value, str):
value = cls(value)
if value == VideoContainer.MP4 or value == VideoContainer.AUTO:
return "mp4"
return ""
@dataclass
class VideoComponents:
"""
Dataclass representing the components of a video.
"""
images: ImageInput
frame_rate: Fraction
audio: Optional[AudioInput] = None
metadata: Optional[dict] = None
class VideoInput(ABC):
"""
Abstract base class for video input types.
"""
@abstractmethod
def get_components(self) -> VideoComponents:
"""
Abstract method to get the video components (images, audio, and frame rate).
Returns:
VideoComponents containing images, audio, and frame rate
"""
pass
@abstractmethod
def save_to(
self,
path: str,
format: VideoContainer = VideoContainer.AUTO,
codec: VideoCodec = VideoCodec.AUTO,
metadata: Optional[dict] = None
):
"""
Abstract method to save the video input to a file.
"""
pass
# Provide a default implementation, but subclasses can provide optimized versions
# if possible.
def get_dimensions(self) -> tuple[int, int]:
"""
Returns the dimensions of the video input.
Returns:
Tuple of (width, height)
"""
components = self.get_components()
return components.images.shape[2], components.images.shape[1]
class VideoFromFile(VideoInput): class VideoFromFile(VideoInput):
""" """
@ -123,7 +22,7 @@ class VideoFromFile(VideoInput):
Initialize the VideoFromFile object based off of either a path on disk or a BytesIO object Initialize the VideoFromFile object based off of either a path on disk or a BytesIO object
containing the file contents. containing the file contents.
""" """
self.file = file self.__file = file
def get_dimensions(self) -> tuple[int, int]: def get_dimensions(self) -> tuple[int, int]:
""" """
@ -132,14 +31,14 @@ class VideoFromFile(VideoInput):
Returns: Returns:
Tuple of (width, height) Tuple of (width, height)
""" """
if isinstance(self.file, io.BytesIO): if isinstance(self.__file, io.BytesIO):
self.file.seek(0) # Reset the BytesIO object to the beginning self.__file.seek(0) # Reset the BytesIO object to the beginning
with av.open(self.file, mode='r') as container: with av.open(self.__file, mode='r') as container:
for stream in container.streams: for stream in container.streams:
if stream.type == 'video': if stream.type == 'video':
assert isinstance(stream, av.VideoStream) assert isinstance(stream, av.VideoStream)
return stream.width, stream.height return stream.width, stream.height
raise ValueError(f"No video stream found in file '{self.file}'") raise ValueError(f"No video stream found in file '{self.__file}'")
def get_components_internal(self, container: InputContainer) -> VideoComponents: def get_components_internal(self, container: InputContainer) -> VideoComponents:
# Get video frames # Get video frames
@ -182,11 +81,11 @@ class VideoFromFile(VideoInput):
return VideoComponents(images=images, audio=audio, frame_rate=frame_rate, metadata=metadata) return VideoComponents(images=images, audio=audio, frame_rate=frame_rate, metadata=metadata)
def get_components(self) -> VideoComponents: def get_components(self) -> VideoComponents:
if isinstance(self.file, io.BytesIO): if isinstance(self.__file, io.BytesIO):
self.file.seek(0) # Reset the BytesIO object to the beginning self.__file.seek(0) # Reset the BytesIO object to the beginning
with av.open(self.file, mode='r') as container: with av.open(self.__file, mode='r') as container:
return self.get_components_internal(container) return self.get_components_internal(container)
raise ValueError(f"No video stream found in file '{self.file}'") raise ValueError(f"No video stream found in file '{self.__file}'")
def save_to( def save_to(
self, self,
@ -195,9 +94,9 @@ class VideoFromFile(VideoInput):
codec: VideoCodec = VideoCodec.AUTO, codec: VideoCodec = VideoCodec.AUTO,
metadata: Optional[dict] = None metadata: Optional[dict] = None
): ):
if isinstance(self.file, io.BytesIO): if isinstance(self.__file, io.BytesIO):
self.file.seek(0) # Reset the BytesIO object to the beginning self.__file.seek(0) # Reset the BytesIO object to the beginning
with av.open(self.file, mode='r') as container: with av.open(self.__file, mode='r') as container:
container_format = container.format.name container_format = container.format.name
video_encoding = container.streams.video[0].codec.name if len(container.streams.video) > 0 else None video_encoding = container.streams.video[0].codec.name if len(container.streams.video) > 0 else None
reuse_streams = True reuse_streams = True
@ -250,13 +149,13 @@ class VideoFromComponents(VideoInput):
""" """
def __init__(self, components: VideoComponents): def __init__(self, components: VideoComponents):
self.components = components self.__components = components
def get_components(self) -> VideoComponents: def get_components(self) -> VideoComponents:
return VideoComponents( return VideoComponents(
images=self.components.images, images=self.__components.images,
audio=self.components.audio, audio=self.__components.audio,
frame_rate=self.components.frame_rate frame_rate=self.__components.frame_rate
) )
def save_to( def save_to(
@ -276,24 +175,24 @@ class VideoFromComponents(VideoInput):
for key, value in metadata.items(): for key, value in metadata.items():
output.metadata[key] = json.dumps(value) output.metadata[key] = json.dumps(value)
frame_rate = Fraction(round(self.components.frame_rate * 1000), 1000) frame_rate = Fraction(round(self.__components.frame_rate * 1000), 1000)
# Create a video stream # Create a video stream
video_stream = output.add_stream('h264', rate=frame_rate) video_stream = output.add_stream('h264', rate=frame_rate)
video_stream.width = self.components.images.shape[2] video_stream.width = self.__components.images.shape[2]
video_stream.height = self.components.images.shape[1] video_stream.height = self.__components.images.shape[1]
video_stream.pix_fmt = 'yuv420p' video_stream.pix_fmt = 'yuv420p'
# Create an audio stream # Create an audio stream
audio_sample_rate = 1 audio_sample_rate = 1
audio_stream: Optional[av.AudioStream] = None audio_stream: Optional[av.AudioStream] = None
if self.components.audio: if self.__components.audio:
audio_sample_rate = int(self.components.audio['sample_rate']) audio_sample_rate = int(self.__components.audio['sample_rate'])
audio_stream = output.add_stream('aac', rate=audio_sample_rate) audio_stream = output.add_stream('aac', rate=audio_sample_rate)
audio_stream.sample_rate = audio_sample_rate audio_stream.sample_rate = audio_sample_rate
audio_stream.format = 'fltp' audio_stream.format = 'fltp'
# Encode video # Encode video
for i, frame in enumerate(self.components.images): for i, frame in enumerate(self.__components.images):
img = (frame * 255).clamp(0, 255).byte().cpu().numpy() # shape: (H, W, 3) img = (frame * 255).clamp(0, 255).byte().cpu().numpy() # shape: (H, W, 3)
frame = av.VideoFrame.from_ndarray(img, format='rgb24') frame = av.VideoFrame.from_ndarray(img, format='rgb24')
frame = frame.reformat(format='yuv420p') # Convert to YUV420P as required by h264 frame = frame.reformat(format='yuv420p') # Convert to YUV420P as required by h264
@ -304,15 +203,15 @@ class VideoFromComponents(VideoInput):
packet = video_stream.encode(None) packet = video_stream.encode(None)
output.mux(packet) output.mux(packet)
if audio_stream and self.components.audio: if audio_stream and self.__components.audio:
# Encode audio # Encode audio
samples_per_frame = int(audio_sample_rate / frame_rate) samples_per_frame = int(audio_sample_rate / frame_rate)
num_frames = self.components.audio['waveform'].shape[2] // samples_per_frame num_frames = self.__components.audio['waveform'].shape[2] // samples_per_frame
for i in range(num_frames): for i in range(num_frames):
start = i * samples_per_frame start = i * samples_per_frame
end = start + samples_per_frame end = start + samples_per_frame
# TODO(Feature) - Add support for stereo audio # TODO(Feature) - Add support for stereo audio
chunk = self.components.audio['waveform'][0, 0, start:end].unsqueeze(0).numpy() chunk = self.__components.audio['waveform'][0, 0, start:end].unsqueeze(0).numpy()
audio_frame = av.AudioFrame.from_ndarray(chunk, format='fltp', layout='mono') audio_frame = av.AudioFrame.from_ndarray(chunk, format='fltp', layout='mono')
audio_frame.sample_rate = audio_sample_rate audio_frame.sample_rate = audio_sample_rate
audio_frame.pts = i * samples_per_frame audio_frame.pts = i * samples_per_frame

View File

@ -0,0 +1,8 @@
from .video_types import VideoContainer, VideoCodec, VideoComponents
__all__ = [
# Utility Types
"VideoContainer",
"VideoCodec",
"VideoComponents",
]

View File

@ -0,0 +1,51 @@
from __future__ import annotations
from dataclasses import dataclass
from enum import Enum
from fractions import Fraction
from typing import Optional
from comfy_api.input import ImageInput, AudioInput
class VideoCodec(str, Enum):
AUTO = "auto"
H264 = "h264"
@classmethod
def as_input(cls) -> list[str]:
"""
Returns a list of codec names that can be used as node input.
"""
return [member.value for member in cls]
class VideoContainer(str, Enum):
AUTO = "auto"
MP4 = "mp4"
@classmethod
def as_input(cls) -> list[str]:
"""
Returns a list of container names that can be used as node input.
"""
return [member.value for member in cls]
@classmethod
def get_extension(cls, value) -> str:
"""
Returns the file extension for the container.
"""
if isinstance(value, str):
value = cls(value)
if value == VideoContainer.MP4 or value == VideoContainer.AUTO:
return "mp4"
return ""
@dataclass
class VideoComponents:
"""
Dataclass representing the components of a video.
"""
images: ImageInput
frame_rate: Fraction
audio: Optional[AudioInput] = None
metadata: Optional[dict] = None

View File

@ -7,7 +7,10 @@ import folder_paths
import json import json
from typing import Optional, Literal from typing import Optional, Literal
from fractions import Fraction from fractions import Fraction
from comfy.comfy_types import IO, FileLocator, ComfyNodeABC, VideoInput, AudioInput, ImageInput, VideoFromComponents, VideoContainer, VideoCodec, VideoComponents, VideoFromFile from comfy.comfy_types import IO, FileLocator, ComfyNodeABC
from comfy_api.input import ImageInput, AudioInput, VideoInput
from comfy_api.util import VideoContainer, VideoCodec, VideoComponents
from comfy_api.input_impl import VideoFromFile, VideoFromComponents
from comfy.cli_args import args from comfy.cli_args import args
class SaveWEBM: class SaveWEBM: