Zero2x / app.py
calebmills99's picture
Rename App.py to app.py
02f9f3d verified
Raw
History Blame Contribute Delete
26.1 kB
from spandrel import ModelLoader
import torch
from pathlib import Path
import gradio as App
import logging
import time
import cv2
import os
# Conditional import for Hugging Face Spaces
try:
import spaces
HF_SPACES_AVAILABLE = True
except ImportError:
# Create a dummy decorator for local development
class DummySpaces:
@staticmethod
def GPU(func):
return func
spaces = DummySpaces()
HF_SPACES_AVAILABLE = True
from gradio import themes
from rich.console import Console
from rich.logging import RichHandler
from Scripts.SAD import GetDifferenceRectangles
from Scripts.ORB import DetectMotionWithOrb
# ============================== #
# Core Settings #
# ============================== #
# Use default theme for HF compatibility
Theme = None # Will use Gradio's default theme
ModelDir = Path("./Models")
TempDir = Path("./Temp")
os.environ["GRADIO_TEMP_DIR"] = str(TempDir)
ModelFileType = ".pth"
# ============================== #
# Logging #
# ============================== #
logging.basicConfig(
level=logging.INFO,
format="%(message)s",
datefmt="[%X]",
handlers=[
RichHandler(
console=Console(),
rich_tracebacks=True,
omit_repeated_times=False,
markup=True,
show_path=False,
)
],
)
Logger = logging.getLogger("Zero2x")
logging.getLogger("httpx").setLevel(logging.WARNING)
# ============================== #
# Device Configuration #
# ============================== #
@spaces.GPU
def GetDeviceName():
Device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
Logger.info(f"๐Ÿงช Using device: {str(Device).upper()}")
return Device
Device = GetDeviceName()
# ============================== #
# Utility Functions #
# ============================== #
def HumanizeSeconds(Seconds):
Hours = int(Seconds // 3600)
Minutes = int((Seconds % 3600) // 60)
Seconds = int(Seconds % 60)
if Hours > 0:
return f"{Hours}h {Minutes}m {Seconds}s"
elif Minutes > 0:
return f"{Minutes}m {Seconds}s"
else:
return f"{Seconds}s"
def HumanizedBytes(Size):
Units = ["B", "KB", "MB", "GB", "TB"]
Index = 0
while Size >= 1024 and Index < len(Units) - 1:
Size /= 1024.0
Index += 1
return f"{Size:.2f} {Units[Index]}"
# ============================== #
# Main Processing Logic #
# ============================== #
class Upscaler:
def __init__(self, preload=False):
self.loaded_models = {}
if preload:
self.preload_models()
else:
Logger.info(
f"๐Ÿ“š Found {len(self.ListModelsFromDisk())} Models In Directory (lazy loading enabled)"
)
def preload_models(self):
"""Preload all models during initialization for faster access"""
model_files = sorted(
[File.stem for File in ModelDir.glob("*" + ModelFileType) if File.is_file()]
)
Logger.info(f"๐Ÿ“š Found {len(model_files)} Models In Directory")
Logger.info("๐Ÿš€ Preloading all models for faster access...")
for i, model_name in enumerate(model_files):
try:
Logger.info(f"๐Ÿ“ฆ Loading model {i+1}/{len(model_files)}: {model_name}")
model = (
ModelLoader()
.load_from_file(ModelDir / (model_name + ModelFileType))
.to(Device)
.eval()
)
self.loaded_models[model_name] = model
Logger.info(f"โœ… Loaded {model_name}")
except Exception as e:
Logger.error(f"โŒ Failed to load {model_name}: {str(e)}")
Logger.info(f"๐ŸŽ‰ Preloaded {len(self.loaded_models)} models successfully!")
def ListModelsFromDisk(self):
"""Get list of model files from disk"""
return sorted(
[File.stem for File in ModelDir.glob("*" + ModelFileType) if File.is_file()]
)
def ListModels(self):
"""Get list of available models (preloaded or from disk)"""
if self.loaded_models:
return list(self.loaded_models.keys())
else:
return self.ListModelsFromDisk()
def LoadModel(self, ModelName):
"""Get preloaded model - much faster than loading from disk"""
if ModelName in self.loaded_models:
Logger.info(f"โšก Using preloaded model: {ModelName}")
return self.loaded_models[ModelName]
else:
Logger.warning(f"๐Ÿ”„ Model {ModelName} not preloaded, loading from disk...")
try:
torch.cuda.empty_cache()
Model = (
ModelLoader()
.load_from_file(ModelDir / (ModelName + ModelFileType))
.to(Device)
.eval()
)
self.loaded_models[ModelName] = Model
Logger.info(f"๐Ÿค– Loaded Model {ModelName} Onto {str(Device).upper()}")
return Model
except Exception as e:
Logger.error(f"โŒ Failed to load {ModelName}: {str(e)}")
Logger.error(
f"๐Ÿ’ก Model file may be corrupted. Try downloading a fresh copy."
)
raise ValueError(f"Model {ModelName} failed to load: {str(e)}")
def UnloadModel(self):
if Device.type == "cuda":
torch.cuda.empty_cache()
Logger.info("๐Ÿค– Model Unloaded Successfully")
def CleanUp(self):
self.UnloadModel()
Logger.info("๐Ÿงน Temporary Files Cleaned Up")
@spaces.GPU
def UpscaleFullFrame(self, Model, Frame):
FrameRgb = cv2.cvtColor(Frame, cv2.COLOR_BGR2RGB)
FrameForTorch = FrameRgb.transpose(2, 0, 1)
FrameForTorch = (
torch.from_numpy(FrameForTorch).unsqueeze(0).to(Device).float() / 255.0
)
OutputFrame = Model(FrameForTorch)[0].cpu().numpy().transpose(1, 2, 0) * 255.0
OutputFrame = cv2.cvtColor(OutputFrame.astype("uint8"), cv2.COLOR_RGB2BGR)
return OutputFrame
@spaces.GPU
def UpscaleRegions(
self,
Model,
Frame,
PrevFrame,
UpscaledPrevFrame,
InputThreshold,
InputMinPercentage,
InputMaxRectangles,
InputPadding,
InputSegmentRows,
InputSegmentColumns,
):
DiffResult = GetDifferenceRectangles(
PrevFrame,
Frame,
Threshold=InputThreshold,
Rows=InputSegmentRows,
Columns=InputSegmentColumns,
Padding=InputPadding,
)
SimilarityPercentage = DiffResult["SimilarPercentage"]
Rectangles = DiffResult["Rectangles"]
Cols = DiffResult["Columns"]
Rows = DiffResult["Rows"]
FrameHeight, FrameWidth = Frame.shape[:2]
SegmentWidth = FrameWidth // Cols
SegmentHeight = FrameHeight // Rows
UseRegions = False
RegionLog = "๐ŸŸฅ"
if (
SimilarityPercentage > InputMinPercentage
and len(Rectangles) < InputMaxRectangles
):
UpscaleFactorY = UpscaledPrevFrame.shape[0] // FrameHeight
UpscaleFactorX = UpscaledPrevFrame.shape[1] // FrameWidth
OutputFrame = UpscaledPrevFrame.copy()
for X, Y, W, H in Rectangles:
X1 = X * SegmentWidth
Y1 = Y * SegmentHeight
X2 = FrameWidth if X + W == Cols else X1 + W * SegmentWidth
Y2 = FrameHeight if Y + H == Rows else Y1 + H * SegmentHeight
Region = Frame[Y1:Y2, X1:X2]
RegionRgb = cv2.cvtColor(Region, cv2.COLOR_BGR2RGB)
RegionTorch = (
torch.from_numpy(RegionRgb.transpose(2, 0, 1))
.unsqueeze(0)
.to(Device)
.float()
/ 255.0
)
UpscaledRegion = (
Model(RegionTorch)[0].cpu().numpy().transpose(1, 2, 0) * 255.0
)
UpscaledRegion = cv2.cvtColor(
UpscaledRegion.astype("uint8"), cv2.COLOR_RGB2BGR
)
RegionHeight, RegionWidth = Region.shape[:2]
UpscaledRegion = cv2.resize(
UpscaledRegion,
(RegionWidth * UpscaleFactorX, RegionHeight * UpscaleFactorY),
interpolation=cv2.INTER_CUBIC,
)
UX1 = X1 * UpscaleFactorX
UY1 = Y1 * UpscaleFactorY
UX2 = UX1 + UpscaledRegion.shape[1]
UY2 = UY1 + UpscaledRegion.shape[0]
OutputFrame[UY1:UY2, UX1:UX2] = UpscaledRegion
RegionLog = "๐ŸŸฉ"
UseRegions = True
else:
OutputFrame = self.UpscaleFullFrame(Model, Frame)
return OutputFrame, SimilarityPercentage, Rectangles, RegionLog, UseRegions
@spaces.GPU
def Process(
self,
InputVideo,
InputModel,
InputUseRegions,
InputThreshold,
InputMinPercentage,
InputMaxRectangles,
InputPadding,
InputSegmentRows,
InputSegmentColumns,
InputFullFrameInterval,
InputMotionThreshold,
Progress=App.Progress(),
):
try:
if not InputVideo:
Logger.warning('โŒ No Video Provided')
App.Warning('โŒ No Video Provided')
return None, None
if not InputModel:
Logger.warning('โŒ No Model Selected')
App.Warning('โŒ No Model Selected - Please add .pth model files to the Models directory')
return None, None
Progress(0, desc='โš™๏ธ Loading Model')
Model = self.LoadModel(InputModel)
Logger.info(f'๐Ÿ“ผ Processing Video: {Path(InputVideo).name}')
Progress(0, desc='๐Ÿ“ผ Processing Video')
Video = cv2.VideoCapture(InputVideo)
if not Video.isOpened():
Logger.error('โŒ Failed to open input video')
return None, None
FrameRate = Video.get(cv2.CAP_PROP_FPS)
FrameCount = int(Video.get(cv2.CAP_PROP_FRAME_COUNT))
Width = int(Video.get(cv2.CAP_PROP_FRAME_WIDTH))
Height = int(Video.get(cv2.CAP_PROP_FRAME_HEIGHT))
if FrameCount <= 0:
Logger.error('โŒ Invalid video: no frames detected')
Video.release()
return None, None
Logger.info(f'๐Ÿ“ Video Properties: {FrameCount} Frames, {FrameRate} FPS, {Width}x{Height}')
PerFrameProgress = 1 / FrameCount
FrameProgress = 0.0
StartTime = time.time()
Times = []
CurrentFrameIndex = 0
PrevFrame = None
UpscaledPrevFrame = None
PartialUpscaleCount = 0
while True:
Ret, Frame = Video.read()
if not Ret:
break
CurrentFrameIndex += 1
# ... (rest of the frame processing loop remains the same)
ForceFull = False
if CurrentFrameIndex == 1 or not InputUseRegions or PartialUpscaleCount >= InputFullFrameInterval:
ForceFull = True
PartialUpscaleCount = 0
if PrevFrame is not None:
IsMotion, _, _ = DetectMotionWithOrb(PrevFrame, Frame, InputMotionThreshold)
if IsMotion:
ForceFull = True
PartialUpscaleCount = 0
Logger.info(f'๐ŸŸจ Frame {CurrentFrameIndex}: Motion Detected - Upscaling Full Frame')
if not ForceFull and PrevFrame is not None and UpscaledPrevFrame is not None:
DiffResult = GetDifferenceRectangles(PrevFrame, Frame, Threshold=InputThreshold, Rows=InputSegmentRows, Columns=InputSegmentColumns, Padding=InputPadding)
if DiffResult['SimilarPercentage'] == 100:
OutputFrame = UpscaledPrevFrame.copy()
Logger.info(f'๐ŸŸฆ Frame {CurrentFrameIndex}: 100% Similar - Copied Previous Upscaled Frame')
cv2.imwrite(f'{TempDir}/Upscaled_Frame_{CurrentFrameIndex:05d}.png', OutputFrame)
PrevFrame = Frame.copy()
UpscaledPrevFrame = OutputFrame.copy()
continue
if ForceFull:
OutputFrame = self.UpscaleFullFrame(Model, Frame)
UseRegions = False
else:
OutputFrame, _, _, _, UseRegions = self.UpscaleRegions(Model, Frame, PrevFrame, UpscaledPrevFrame, InputThreshold, InputMinPercentage, InputMaxRectangles, InputPadding, InputSegmentRows, InputSegmentColumns)
if UseRegions:
PartialUpscaleCount += 1
else:
PartialUpscaleCount = 0
cv2.imwrite(f'{TempDir}/Upscaled_Frame_{CurrentFrameIndex:05d}.png', OutputFrame)
Progress(CurrentFrameIndex / FrameCount, desc=f'๐Ÿ“ฆ Processed Frame {CurrentFrameIndex}/{FrameCount}')
PrevFrame = Frame.copy()
UpscaledPrevFrame = OutputFrame.copy()
Video.release()
Progress(1, desc='๐Ÿ“ฆ Creating Final Video...')
if CurrentFrameIndex == 0:
Logger.error('โŒ No frames were processed')
return None, None
FirstFramePath = f'{TempDir}/Upscaled_Frame_00001.png'
if not os.path.exists(FirstFramePath):
Logger.error('โŒ No processed frames found')
return None, None
FirstFrame = cv2.imread(FirstFramePath)
if FirstFrame is None:
Logger.error('โŒ Could not read first processed frame')
return None, None
OutputHeight, OutputWidth = FirstFrame.shape[:2]
InputPath = Path(InputVideo)
OutputPath = TempDir / f'Upscaled_{InputPath.stem}_{InputModel}.mp4'
codecs_to_try = [
('mp4v', cv2.VideoWriter_fourcc('m', 'p', '4', 'v')),
('avc1', cv2.VideoWriter_fourcc('a', 'v', 'c', '1')),
('H264', cv2.VideoWriter_fourcc('H', '2', '6', '4')),
]
VideoWriter = None
successful_codec = None
for codec_name, fourcc in codecs_to_try:
VideoWriter = cv2.VideoWriter(str(OutputPath), fourcc, FrameRate, (OutputWidth, OutputHeight))
if VideoWriter.isOpened():
successful_codec = codec_name
Logger.info(f'โœ… Video writer initialized with codec: {codec_name}')
break
VideoWriter.release()
VideoWriter = None
if VideoWriter is None:
Logger.error('โŒ Failed to initialize video writer with any codec')
return None, None
frames_written = 0
for i in range(1, CurrentFrameIndex + 1):
FramePath = f'{TempDir}/Upscaled_Frame_{i:05d}.png'
if os.path.exists(FramePath):
Frame = cv2.imread(FramePath)
if Frame is not None:
if Frame.shape[:2] != (OutputHeight, OutputWidth):
Frame = cv2.resize(Frame, (OutputWidth, OutputHeight))
VideoWriter.write(Frame)
frames_written += 1
try:
os.remove(FramePath)
except Exception as e:
Logger.warning(f'โš ๏ธ Could not remove temp frame {FramePath}: {e}')
VideoWriter.release()
if frames_written == 0:
Logger.error('โŒ No frames were written to output video')
return None, None
Logger.info(f'โœ… Output video created: {OutputPath}')
return str(OutputPath), str(OutputPath)
finally:
self.CleanUp()
# ============================== #
# Global Upscaler Instance #
# ============================== #
# Create global upscaler instance with lazy loading for faster startup
upscaler = Upscaler(preload=False)
# ============================== #
# Simplified UI #
# ============================== #
with App.Blocks(title="Zero2x Video Upscaler", delete_cache=(-1, 1800)) as Interface:
App.Markdown("# ๐ŸŽž๏ธ Zero2x Video Upscaler")
# Check if models are available
try:
ModelNames = upscaler.ListModels()
if not ModelNames:
ModelNames = ["test_model"] # Fallback for testing
except Exception as e:
Logger.warning(f"Error loading models: {e}")
ModelNames = ["test_model"] # Fallback for testing
# Main Video Upscaling Interface
if not ModelNames or ModelNames == ["test_model"]:
App.Markdown(
"""
## โš ๏ธ No Models Found
**No upscaling models (.pth files) were found in the Models directory.**
To use this application, you need to:
1. Download upscaling models (Real-ESRGAN, CUGAN, etc.)
2. Place .pth model files in the `Models/` directory
3. Restart the application
"""
)
with App.Row():
with App.Column():
with App.Group():
InputVideo = App.Video(
label="Input Video", sources=["upload"], height=300
)
if ModelNames:
InputModel = App.Dropdown(
choices=ModelNames,
label="Select Model",
value=ModelNames[0] if ModelNames else None,
)
else:
InputModel = App.Dropdown(
choices=[],
label="Select Model (No models found)",
value=None,
interactive=False,
)
with App.Accordion(label="โš™๏ธ Advanced Settings", open=False):
with App.Group():
InputUseRegions = App.Checkbox(
label="Use Regions",
value=False,
info="Use regions to upscale only the different parts of the video (โšก๏ธ Experimental, Faster)",
interactive=bool(ModelNames),
)
InputThreshold = App.Slider(
label="Threshold",
value=2,
minimum=0,
maximum=10,
step=0.5,
info="Threshold for the SAD algorithm to detect different regions",
interactive=False,
)
InputPadding = App.Slider(
label="Padding",
value=1,
minimum=0,
maximum=5,
step=1,
info="Extra padding to include neighboring pixels in the SAD algorithm",
interactive=False,
)
InputMinPercentage = App.Slider(
label="Min Percentage",
value=50,
minimum=0,
maximum=100,
step=1,
info="Minimum percentage of similarity to consider upscaling the full frame",
interactive=False,
)
InputMaxRectangles = App.Slider(
label="Max Rectangles",
value=10,
minimum=1,
maximum=16,
step=1,
info="Maximum number of rectangles to consider upscaling the full frame",
interactive=False,
)
with App.Row():
InputSegmentRows = App.Slider(
label="Segment Rows",
value=32,
minimum=1,
maximum=64,
step=1,
info="Number of rows to segment the video into for processing",
interactive=False,
)
InputSegmentColumns = App.Slider(
label="Segment Columns",
value=48,
minimum=1,
maximum=64,
step=1,
info="Number of columns to segment the video into for processing",
interactive=False,
)
InputFullFrameInterval = App.Slider(
label="Full Frame Interval",
value=5,
minimum=1,
maximum=100,
step=1,
info="Force a full-frame upscale every N frames (set to 1 to always upscale full frame)",
interactive=False,
)
InputMotionThreshold = App.Slider(
label="Motion Threshold",
value=1,
minimum=0,
maximum=10,
step=0.5,
info="Threshold for the motion detection algorithm to consider a frame as different",
interactive=False,
)
if ModelNames and ModelNames != ["test_model"]:
SubmitButton = App.Button("๐Ÿš€ Upscale Video")
else:
SubmitButton = App.Button("โŒ No Models Available", interactive=False)
with App.Column(show_progress=True):
with App.Group():
OutputVideo = App.Video(
label="Output Video",
height=300,
interactive=False,
)
OutputDownload = App.DownloadButton(
label="๐Ÿ’พ Download Video", interactive=False
)
# About Section
with App.Accordion(label="โ„น๏ธ About Zero2x", open=False):
App.Markdown(
"""
## โœจ Zero2x Video Upscaler
**Zero2x** is an advanced video upscaling tool using AI models:
- ๐ŸŽฌ **AI Video Upscaling** using multiple deep learning models
- ๐Ÿง‘โ€๐Ÿ”ฌ **Advanced Techniques**: SAD algorithm for region detection, ORB motion detection
- ๐Ÿš€ **GPU Acceleration**: CUDA support with automatic memory management
### ๐Ÿš€ Getting Started
1. **Add Models**: Place .pth model files in the `WorkingModels/` directory
2. **Upload Video**: Choose a video file to upscale
3. **Configure Settings**: Adjust parameters as needed
4. **Process**: Let the AI enhance your video!
"""
)
# Event Handlers
def toggle_region_inputs(use_regions):
# This is a more explicit way to update the interactivity of the sliders
# instead of relying on a magic number in a loop.
return {
InputThreshold: App.update(interactive=use_regions),
InputMinPercentage: App.update(interactive=use_regions),
InputMaxRectangles: App.update(interactive=use_regions),
InputPadding: App.update(interactive=use_regions),
InputSegmentRows: App.update(interactive=use_regions),
InputSegmentColumns: App.update(interactive=use_regions),
InputFullFrameInterval: App.update(interactive=use_regions),
InputMotionThreshold: App.update(interactive=use_regions),
}
# Wire up events
InputUseRegions.change(
fn=toggle_region_inputs,
inputs=[InputUseRegions],
outputs=[
InputThreshold,
InputMinPercentage,
InputMaxRectangles,
InputPadding,
InputSegmentRows,
InputSegmentColumns,
InputFullFrameInterval,
InputMotionThreshold,
],
)
if ModelNames and ModelNames != ["test_model"] and upscaler:
SubmitButton.click(
fn=upscaler.Process,
inputs=[
InputVideo,
InputModel,
InputUseRegions,
InputThreshold,
InputMinPercentage,
InputMaxRectangles,
InputPadding,
InputSegmentRows,
InputSegmentColumns,
InputFullFrameInterval,
InputMotionThreshold,
],
outputs=[OutputVideo, OutputDownload],
)
if __name__ == "__main__":
os.makedirs(ModelDir, exist_ok=True)
os.makedirs(TempDir, exist_ok=True)
Logger.info("๐Ÿš€ Starting Zero2x Video Upscaler")
# Launch configuration for Hugging Face deployment
Interface.launch(server_name="0.0.0.0", server_port=7860, share=True)