File: //opt/SkyReels-V2/torch_wrapper.py
#!/usr/bin/env python3
"""
PyTorch-Wrapper für SkyReels-V2
Dieser Wrapper isoliert PyTorch-Importe, um Konflikte mit Streamlit zu vermeiden
"""
import os
import sys
import importlib
# Erkenne, ob wir in Streamlit ausgeführt werden
IN_STREAMLIT = 'streamlit' in sys.modules
# PyTorch-Module
_torch = None
_diffusers = None
_transformers = None
_safetensors = None
def load_pytorch():
"""Lädt alle PyTorch-bezogenen Module"""
global _torch, _diffusers, _transformers, _safetensors
# Importiere PyTorch
if _torch is None:
try:
import torch
_torch = torch
except ImportError:
print("PyTorch konnte nicht importiert werden. Bitte installieren Sie es mit: pip install torch")
return False
# Importiere Diffusers
if _diffusers is None:
try:
import diffusers
_diffusers = diffusers
except ImportError:
print("Diffusers konnte nicht importiert werden. Bitte installieren Sie es mit: pip install diffusers")
return False
# Importiere Transformers
if _transformers is None:
try:
import transformers
_transformers = transformers
except ImportError:
print("Transformers konnte nicht importiert werden. Bitte installieren Sie es mit: pip install transformers")
return False
# Importiere Safetensors
if _safetensors is None:
try:
import safetensors
_safetensors = safetensors
except ImportError:
print("Safetensors konnte nicht importiert werden. Bitte installieren Sie es mit: pip install safetensors")
return False
return True
def check_dependencies():
"""Prüft, ob alle Abhängigkeiten installiert sind"""
missing = []
try:
import torch
except ImportError:
missing.append("torch")
try:
import diffusers
except ImportError:
missing.append("diffusers")
try:
import transformers
except ImportError:
missing.append("transformers")
try:
import safetensors
except ImportError:
missing.append("safetensors")
return missing
def download_model(model_id):
"""
Lädt ein Modell als separater Prozess herunter, um Streamlit-Konflikte zu vermeiden
"""
# Füge das aktuelle Verzeichnis zum PYTHONPATH hinzu
current_dir = os.path.dirname(os.path.abspath(__file__))
if current_dir not in sys.path:
sys.path.append(current_dir)
if not load_pytorch():
return False, "Konnte PyTorch-Module nicht laden"
try:
from skyreels_v2_infer.modules import download_model as _download_model
return True, _download_model(model_id)
except Exception as e:
return False, str(e)
if __name__ == "__main__":
# Wenn direkt ausgeführt, prüfe die Abhängigkeiten
missing = check_dependencies()
if missing:
print(f"Fehlende Abhängigkeiten: {', '.join(missing)}")
print("Bitte installieren Sie diese mit: pip install " + " ".join(missing))
sys.exit(1)
else:
print("Alle Abhängigkeiten sind installiert!")
# Teste den Import der Module
if load_pytorch():
print("PyTorch und verwandte Module erfolgreich geladen!")
else:
print("Fehler beim Laden von PyTorch und verwandten Modulen.")
sys.exit(1)
sys.exit(0)