Add extensive debug logging to DisTorch implementation
- Added comprehensive file logging with timestamps - Log file created in logs/distorch_TIMESTAMP.log - Debug logging for all major operations: - Model structure analysis - Memory calculations - Block partitioning strategy - Device movements and swaps - Hook installations - Performance metrics - Console output for important INFO messages - Full exception tracebacks captured
This commit is contained in:
+267
-67
@@ -8,6 +8,49 @@ import logging
|
||||
from typing import Dict, List, Tuple, Optional, Any
|
||||
from dataclasses import dataclass
|
||||
import gc
|
||||
import os
|
||||
from datetime import datetime
|
||||
import traceback
|
||||
import json
|
||||
|
||||
# Set up file logging for DisTorch
|
||||
log_dir = os.path.join(os.path.dirname(os.path.dirname(__file__)), "logs")
|
||||
os.makedirs(log_dir, exist_ok=True)
|
||||
log_file = os.path.join(log_dir, f"distorch_{datetime.now().strftime('%Y%m%d_%H%M%S')}.log")
|
||||
|
||||
# Configure file handler
|
||||
file_handler = logging.FileHandler(log_file, mode='w')
|
||||
file_handler.setLevel(logging.DEBUG)
|
||||
file_formatter = logging.Formatter(
|
||||
'%(asctime)s.%(msecs)03d - [%(name)s] - %(levelname)s - %(funcName)s:%(lineno)d - %(message)s',
|
||||
datefmt='%Y-%m-%d %H:%M:%S'
|
||||
)
|
||||
file_handler.setFormatter(file_formatter)
|
||||
|
||||
# Create DisTorch logger
|
||||
distorch_logger = logging.getLogger("DisTorch")
|
||||
distorch_logger.setLevel(logging.DEBUG)
|
||||
distorch_logger.addHandler(file_handler)
|
||||
|
||||
# Also add console handler for important messages
|
||||
console_handler = logging.StreamHandler()
|
||||
console_handler.setLevel(logging.INFO)
|
||||
console_formatter = logging.Formatter('[DisTorch] %(levelname)s: %(message)s')
|
||||
console_handler.setFormatter(console_formatter)
|
||||
distorch_logger.addHandler(console_handler)
|
||||
|
||||
distorch_logger.info(f"DisTorch logging initialized. Log file: {log_file}")
|
||||
distorch_logger.debug("="*80)
|
||||
distorch_logger.debug("DISTORCH BLOCK SWAP MODULE LOADED")
|
||||
distorch_logger.debug(f"PyTorch version: {torch.__version__}")
|
||||
distorch_logger.debug(f"CUDA available: {torch.cuda.is_available()}")
|
||||
if torch.cuda.is_available():
|
||||
distorch_logger.debug(f"CUDA device count: {torch.cuda.device_count()}")
|
||||
for i in range(torch.cuda.device_count()):
|
||||
distorch_logger.debug(f" Device {i}: {torch.cuda.get_device_name(i)}")
|
||||
distorch_logger.debug(f" Memory: {torch.cuda.get_device_properties(i).total_memory / (1024**3):.2f} GB")
|
||||
distorch_logger.debug("="*80)
|
||||
|
||||
|
||||
@dataclass
|
||||
class BlockSwapConfig:
|
||||
@@ -19,36 +62,99 @@ class BlockSwapConfig:
|
||||
use_non_blocking: bool = False # Non-blocking transfers
|
||||
|
||||
def __post_init__(self):
|
||||
self.swap_device = torch.device(self.swap_device)
|
||||
self.compute_device = torch.device(self.compute_device)
|
||||
distorch_logger.debug(f"BlockSwapConfig.__post_init__ called")
|
||||
distorch_logger.debug(f" Raw swap_device: {self.swap_device}")
|
||||
distorch_logger.debug(f" Raw compute_device: {self.compute_device}")
|
||||
distorch_logger.debug(f" virtual_vram_gb: {self.virtual_vram_gb}")
|
||||
distorch_logger.debug(f" swap_space_gb: {self.swap_space_gb}")
|
||||
distorch_logger.debug(f" use_non_blocking: {self.use_non_blocking}")
|
||||
|
||||
try:
|
||||
self.swap_device = torch.device(self.swap_device)
|
||||
self.compute_device = torch.device(self.compute_device)
|
||||
distorch_logger.debug(f" Converted swap_device: {self.swap_device}")
|
||||
distorch_logger.debug(f" Converted compute_device: {self.compute_device}")
|
||||
except Exception as e:
|
||||
distorch_logger.error(f"Error converting devices: {e}")
|
||||
distorch_logger.error(traceback.format_exc())
|
||||
raise
|
||||
|
||||
|
||||
class BlockSwapManager:
|
||||
"""Manages block swapping for transformer models"""
|
||||
|
||||
def __init__(self, model: torch.nn.Module, config: BlockSwapConfig):
|
||||
distorch_logger.debug("BlockSwapManager.__init__ called")
|
||||
distorch_logger.debug(f" Model type: {type(model).__name__}")
|
||||
distorch_logger.debug(f" Model device: {next(model.parameters()).device if any(model.parameters()) else 'no params'}")
|
||||
|
||||
self.model = model
|
||||
self.config = config
|
||||
self.blocks = []
|
||||
self.current_block_idx = -1
|
||||
self.hooks = []
|
||||
self.swap_count = 0 # Track number of swaps for debugging
|
||||
|
||||
# Calculate model size
|
||||
distorch_logger.debug("Calculating model size...")
|
||||
self.model_size_gb = self._calculate_model_size()
|
||||
logging.info(f"[BlockSwap] Model size: {self.model_size_gb:.2f} GB")
|
||||
distorch_logger.info(f"Model size: {self.model_size_gb:.2f} GB")
|
||||
|
||||
# Log model structure
|
||||
self._log_model_structure()
|
||||
|
||||
# Partition model into blocks
|
||||
distorch_logger.debug("Partitioning model into blocks...")
|
||||
self._partition_model()
|
||||
|
||||
# Install hooks
|
||||
distorch_logger.debug("Installing forward hooks...")
|
||||
self._install_hooks()
|
||||
|
||||
distorch_logger.debug("BlockSwapManager initialization complete")
|
||||
|
||||
def _log_model_structure(self):
|
||||
"""Log the model structure for debugging"""
|
||||
distorch_logger.debug("Model structure analysis:")
|
||||
module_count = 0
|
||||
param_count = 0
|
||||
module_types = {}
|
||||
|
||||
for name, module in self.model.named_modules():
|
||||
module_count += 1
|
||||
module_type = type(module).__name__
|
||||
module_types[module_type] = module_types.get(module_type, 0) + 1
|
||||
|
||||
# Count parameters in this module
|
||||
module_params = sum(p.numel() for p in module.parameters(recurse=False))
|
||||
if module_params > 0:
|
||||
param_count += module_params
|
||||
if module_count <= 10: # Log first 10 modules with params
|
||||
distorch_logger.debug(f" {name}: {module_type} ({module_params:,} params)")
|
||||
|
||||
distorch_logger.debug(f"Total modules: {module_count}")
|
||||
distorch_logger.debug(f"Total parameters: {param_count:,}")
|
||||
distorch_logger.debug("Module type distribution:")
|
||||
for module_type, count in sorted(module_types.items(), key=lambda x: x[1], reverse=True)[:10]:
|
||||
distorch_logger.debug(f" {module_type}: {count}")
|
||||
|
||||
def _calculate_model_size(self) -> float:
|
||||
"""Calculate total model size in GB"""
|
||||
total_bytes = 0
|
||||
for param in self.model.parameters():
|
||||
param_count = 0
|
||||
|
||||
for name, param in self.model.named_parameters():
|
||||
if param.data is not None:
|
||||
total_bytes += param.element_size() * param.nelement()
|
||||
param_bytes = param.element_size() * param.nelement()
|
||||
total_bytes += param_bytes
|
||||
param_count += 1
|
||||
|
||||
if param_count <= 5: # Log first 5 parameters
|
||||
distorch_logger.debug(f" Param {name}: shape={param.shape}, bytes={param_bytes:,}")
|
||||
|
||||
distorch_logger.debug(f"Total parameters: {param_count}")
|
||||
distorch_logger.debug(f"Total bytes: {total_bytes:,}")
|
||||
|
||||
return total_bytes / (1024**3)
|
||||
|
||||
def _get_module_size(self, module: torch.nn.Module) -> float:
|
||||
@@ -61,6 +167,7 @@ class BlockSwapManager:
|
||||
|
||||
def _partition_model(self):
|
||||
"""Partition model into swappable blocks based on swap_space_gb"""
|
||||
distorch_logger.debug("Starting model partitioning...")
|
||||
|
||||
# Find transformer blocks (common patterns)
|
||||
transformer = None
|
||||
@@ -69,19 +176,30 @@ class BlockSwapManager:
|
||||
# Try to find transformer module
|
||||
for name, module in self.model.named_modules():
|
||||
# Common transformer patterns
|
||||
if any(pattern in name.lower() for pattern in ['transformer', 'diffusion_model', 'unet']):
|
||||
patterns = ['transformer', 'diffusion_model', 'unet', 'encoder', 'decoder']
|
||||
if any(pattern in name.lower() for pattern in patterns):
|
||||
distorch_logger.debug(f"Found potential transformer module: {name} ({type(module).__name__})")
|
||||
|
||||
# Check if it has sequential blocks
|
||||
if hasattr(module, 'blocks') or hasattr(module, 'layers'):
|
||||
if hasattr(module, 'blocks'):
|
||||
transformer = module
|
||||
if hasattr(module, 'blocks'):
|
||||
transformer_blocks = list(module.blocks)
|
||||
elif hasattr(module, 'layers'):
|
||||
transformer_blocks = list(module.layers)
|
||||
transformer_blocks = list(module.blocks)
|
||||
distorch_logger.debug(f" Found {len(transformer_blocks)} blocks in {name}.blocks")
|
||||
break
|
||||
elif hasattr(module, 'layers'):
|
||||
transformer = module
|
||||
transformer_blocks = list(module.layers)
|
||||
distorch_logger.debug(f" Found {len(transformer_blocks)} layers in {name}.layers")
|
||||
break
|
||||
elif hasattr(module, 'transformer_blocks'):
|
||||
transformer = module
|
||||
transformer_blocks = list(module.transformer_blocks)
|
||||
distorch_logger.debug(f" Found {len(transformer_blocks)} transformer_blocks in {name}")
|
||||
break
|
||||
|
||||
if not transformer_blocks:
|
||||
# Fallback: partition all modules
|
||||
logging.warning("[BlockSwap] No transformer blocks found, using fallback partitioning")
|
||||
distorch_logger.warning("No transformer blocks found, using fallback partitioning")
|
||||
self._partition_fallback()
|
||||
return
|
||||
|
||||
@@ -90,12 +208,16 @@ class BlockSwapManager:
|
||||
current_size = 0
|
||||
swap_space_bytes = self.config.swap_space_gb * (1024**3)
|
||||
|
||||
distorch_logger.debug(f"Grouping {len(transformer_blocks)} blocks with swap_space={self.config.swap_space_gb} GB")
|
||||
|
||||
for idx, block in enumerate(transformer_blocks):
|
||||
block_size = self._get_module_size(block) * (1024**3) # Convert to bytes
|
||||
distorch_logger.debug(f" Block {idx}: size={block_size/(1024**3):.3f} GB")
|
||||
|
||||
if current_size + block_size > swap_space_bytes and current_block:
|
||||
# Start new block group
|
||||
self.blocks.append(current_block)
|
||||
distorch_logger.debug(f" Created block group {len(self.blocks)-1} with {len(current_block)} blocks, total size={current_size/(1024**3):.3f} GB")
|
||||
current_block = [block]
|
||||
current_size = block_size
|
||||
else:
|
||||
@@ -105,82 +227,124 @@ class BlockSwapManager:
|
||||
# Add remaining blocks
|
||||
if current_block:
|
||||
self.blocks.append(current_block)
|
||||
distorch_logger.debug(f" Created final block group {len(self.blocks)-1} with {len(current_block)} blocks, total size={current_size/(1024**3):.3f} GB")
|
||||
|
||||
logging.info(f"[BlockSwap] Partitioned into {len(self.blocks)} block groups")
|
||||
distorch_logger.info(f"Partitioned into {len(self.blocks)} block groups")
|
||||
for i, group in enumerate(self.blocks):
|
||||
group_size = sum(self._get_module_size(b) for b in group)
|
||||
logging.info(f" Block group {i}: {len(group)} blocks, {group_size:.2f} GB")
|
||||
distorch_logger.info(f" Block group {i}: {len(group)} blocks, {group_size:.3f} GB")
|
||||
|
||||
def _partition_fallback(self):
|
||||
"""Fallback partitioning when transformer structure is not recognized"""
|
||||
distorch_logger.debug("Using fallback partitioning strategy")
|
||||
all_modules = []
|
||||
|
||||
# Collect all modules with parameters
|
||||
for name, module in self.model.named_modules():
|
||||
if any(param.numel() > 0 for param in module.parameters(recurse=False)):
|
||||
all_modules.append((name, module))
|
||||
param_count = sum(p.numel() for p in module.parameters(recurse=False))
|
||||
if param_count > 0:
|
||||
module_size = self._get_module_size(module)
|
||||
all_modules.append((name, module, module_size))
|
||||
if len(all_modules) <= 10:
|
||||
distorch_logger.debug(f" Module {name}: {param_count:,} params, {module_size:.3f} GB")
|
||||
|
||||
distorch_logger.debug(f"Found {len(all_modules)} modules with parameters")
|
||||
|
||||
# Group by size
|
||||
current_block = []
|
||||
current_size = 0
|
||||
swap_space_bytes = self.config.swap_space_gb * (1024**3)
|
||||
|
||||
for name, module in all_modules:
|
||||
module_size = self._get_module_size(module) * (1024**3)
|
||||
for name, module, module_size_gb in all_modules:
|
||||
module_size = module_size_gb * (1024**3)
|
||||
|
||||
if current_size + module_size > swap_space_bytes and current_block:
|
||||
self.blocks.append([m for _, m in current_block])
|
||||
current_block = [(name, module)]
|
||||
self.blocks.append([m for _, m, _ in current_block])
|
||||
distorch_logger.debug(f" Created block group with {len(current_block)} modules, size={current_size/(1024**3):.3f} GB")
|
||||
current_block = [(name, module, module_size_gb)]
|
||||
current_size = module_size
|
||||
else:
|
||||
current_block.append((name, module))
|
||||
current_block.append((name, module, module_size_gb))
|
||||
current_size += module_size
|
||||
|
||||
if current_block:
|
||||
self.blocks.append([m for _, m in current_block])
|
||||
self.blocks.append([m for _, m, _ in current_block])
|
||||
distorch_logger.debug(f" Created final block group with {len(current_block)} modules, size={current_size/(1024**3):.3f} GB")
|
||||
|
||||
def _install_hooks(self):
|
||||
"""Install forward pre-hooks on blocks"""
|
||||
distorch_logger.debug(f"Installing hooks on {len(self.blocks)} block groups")
|
||||
|
||||
for block_idx, block_group in enumerate(self.blocks):
|
||||
distorch_logger.debug(f" Installing hooks for block group {block_idx} ({len(block_group)} modules)")
|
||||
for module in block_group:
|
||||
hook = module.register_forward_pre_hook(
|
||||
lambda m, i, bidx=block_idx: self._pre_forward_hook(m, i, bidx)
|
||||
)
|
||||
self.hooks.append(hook)
|
||||
|
||||
distorch_logger.debug(f"Installed {len(self.hooks)} hooks total")
|
||||
|
||||
def _pre_forward_hook(self, module: torch.nn.Module, inputs: Tuple, block_idx: int):
|
||||
"""Hook called before forward pass of each block"""
|
||||
if block_idx != self.current_block_idx:
|
||||
distorch_logger.debug(f"Pre-forward hook triggered: current={self.current_block_idx}, needed={block_idx}")
|
||||
self._swap_blocks(self.current_block_idx, block_idx)
|
||||
self.current_block_idx = block_idx
|
||||
return inputs
|
||||
|
||||
def _swap_blocks(self, old_idx: int, new_idx: int):
|
||||
"""Swap blocks between devices"""
|
||||
logging.debug(f"[BlockSwap] Swapping from block {old_idx} to {new_idx}")
|
||||
self.swap_count += 1
|
||||
distorch_logger.debug(f"[Swap #{self.swap_count}] Swapping from block {old_idx} to {new_idx}")
|
||||
distorch_logger.debug(f" Swap device: {self.config.swap_device}")
|
||||
distorch_logger.debug(f" Compute device: {self.config.compute_device}")
|
||||
|
||||
start_time = datetime.now()
|
||||
|
||||
# Offload old block
|
||||
if old_idx >= 0 and old_idx < len(self.blocks):
|
||||
for module in self.blocks[old_idx]:
|
||||
distorch_logger.debug(f" Offloading block {old_idx} to {self.config.swap_device}")
|
||||
for i, module in enumerate(self.blocks[old_idx]):
|
||||
self._move_module(module, self.config.swap_device)
|
||||
if i == 0: # Log first module movement
|
||||
distorch_logger.debug(f" Moved module {type(module).__name__} to {self.config.swap_device}")
|
||||
|
||||
# Load new block
|
||||
if new_idx >= 0 and new_idx < len(self.blocks):
|
||||
for module in self.blocks[new_idx]:
|
||||
distorch_logger.debug(f" Loading block {new_idx} to {self.config.compute_device}")
|
||||
for i, module in enumerate(self.blocks[new_idx]):
|
||||
self._move_module(module, self.config.compute_device)
|
||||
if i == 0: # Log first module movement
|
||||
distorch_logger.debug(f" Moved module {type(module).__name__} to {self.config.compute_device}")
|
||||
|
||||
# Clear cache if needed
|
||||
if self.config.compute_device.type == 'cuda':
|
||||
before_free = torch.cuda.memory_reserved(self.config.compute_device) / (1024**3)
|
||||
torch.cuda.empty_cache()
|
||||
after_free = torch.cuda.memory_reserved(self.config.compute_device) / (1024**3)
|
||||
distorch_logger.debug(f" GPU cache cleared: {before_free:.2f} GB -> {after_free:.2f} GB")
|
||||
|
||||
elapsed = (datetime.now() - start_time).total_seconds()
|
||||
distorch_logger.debug(f" Swap completed in {elapsed:.3f} seconds")
|
||||
|
||||
def _move_module(self, module: torch.nn.Module, device: torch.device):
|
||||
"""Move a module to specified device"""
|
||||
module.to(device, non_blocking=self.config.use_non_blocking)
|
||||
try:
|
||||
module.to(device, non_blocking=self.config.use_non_blocking)
|
||||
except Exception as e:
|
||||
distorch_logger.error(f"Error moving module {type(module).__name__} to {device}: {e}")
|
||||
distorch_logger.error(traceback.format_exc())
|
||||
raise
|
||||
|
||||
def prepare(self):
|
||||
"""Prepare model for inference by moving all blocks to swap device"""
|
||||
logging.info(f"[BlockSwap] Moving all blocks to {self.config.swap_device}")
|
||||
for block_group in self.blocks:
|
||||
distorch_logger.info(f"Preparing model: moving all blocks to {self.config.swap_device}")
|
||||
|
||||
start_time = datetime.now()
|
||||
|
||||
for i, block_group in enumerate(self.blocks):
|
||||
distorch_logger.debug(f" Moving block group {i} ({len(block_group)} modules)")
|
||||
for module in block_group:
|
||||
self._move_module(module, self.config.swap_device)
|
||||
|
||||
@@ -189,15 +353,26 @@ class BlockSwapManager:
|
||||
|
||||
# Clear GPU cache
|
||||
if self.config.compute_device.type == 'cuda':
|
||||
before_free = torch.cuda.memory_reserved(self.config.compute_device) / (1024**3)
|
||||
torch.cuda.empty_cache()
|
||||
after_free = torch.cuda.memory_reserved(self.config.compute_device) / (1024**3)
|
||||
distorch_logger.debug(f"GPU memory after prepare: {before_free:.2f} GB -> {after_free:.2f} GB")
|
||||
|
||||
gc.collect()
|
||||
|
||||
elapsed = (datetime.now() - start_time).total_seconds()
|
||||
distorch_logger.info(f"Model preparation completed in {elapsed:.3f} seconds")
|
||||
|
||||
def cleanup(self):
|
||||
"""Remove hooks and cleanup"""
|
||||
distorch_logger.info("Cleaning up BlockSwapManager")
|
||||
distorch_logger.debug(f" Total swaps performed: {self.swap_count}")
|
||||
|
||||
for hook in self.hooks:
|
||||
hook.remove()
|
||||
self.hooks.clear()
|
||||
logging.info("[BlockSwap] Cleanup complete")
|
||||
|
||||
distorch_logger.info("BlockSwap cleanup complete")
|
||||
|
||||
|
||||
class DisTorch:
|
||||
@@ -208,6 +383,8 @@ class DisTorch:
|
||||
from .. import get_device_list
|
||||
devices = get_device_list()
|
||||
|
||||
distorch_logger.debug(f"DisTorch.INPUT_TYPES called, available devices: {devices}")
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
@@ -250,46 +427,69 @@ class DisTorch:
|
||||
swap_device: str, compute_device: str, use_non_blocking: bool = False):
|
||||
"""Apply block swap configuration to model"""
|
||||
|
||||
logging.info(f"[DisTorch] Configuring block swap:")
|
||||
logging.info(f" Virtual VRAM: {virtual_vram_gb} GB")
|
||||
logging.info(f" Swap space: {swap_space_gb} GB")
|
||||
logging.info(f" Swap device: {swap_device}")
|
||||
logging.info(f" Compute device: {compute_device}")
|
||||
distorch_logger.info("="*60)
|
||||
distorch_logger.info("DisTorch.apply_block_swap called")
|
||||
distorch_logger.info("="*60)
|
||||
|
||||
# Create config
|
||||
config = BlockSwapConfig(
|
||||
virtual_vram_gb=virtual_vram_gb,
|
||||
swap_space_gb=swap_space_gb,
|
||||
swap_device=swap_device,
|
||||
compute_device=compute_device,
|
||||
use_non_blocking=use_non_blocking
|
||||
)
|
||||
distorch_logger.info(f"Configuration:")
|
||||
distorch_logger.info(f" Virtual VRAM: {virtual_vram_gb} GB")
|
||||
distorch_logger.info(f" Swap space: {swap_space_gb} GB")
|
||||
distorch_logger.info(f" Swap device: {swap_device}")
|
||||
distorch_logger.info(f" Compute device: {compute_device}")
|
||||
distorch_logger.info(f" Non-blocking: {use_non_blocking}")
|
||||
|
||||
# Get the actual model (handle ModelPatcher)
|
||||
if hasattr(model, 'model'):
|
||||
actual_model = model.model
|
||||
else:
|
||||
actual_model = model
|
||||
distorch_logger.debug(f"Input model type: {type(model)}")
|
||||
distorch_logger.debug(f"Model attributes: {dir(model)[:10]}...") # Log first 10 attributes
|
||||
|
||||
# Check if model has diffusion_model (common pattern)
|
||||
if hasattr(actual_model, 'diffusion_model'):
|
||||
target_model = actual_model.diffusion_model
|
||||
else:
|
||||
target_model = actual_model
|
||||
|
||||
# Create block swap manager
|
||||
manager = BlockSwapManager(target_model, config)
|
||||
|
||||
# Prepare model (move blocks to swap device)
|
||||
manager.prepare()
|
||||
|
||||
# Store manager on model for later access
|
||||
model._block_swap_manager = manager
|
||||
|
||||
# Also set the load_device attribute if it exists
|
||||
if hasattr(model, 'load_device'):
|
||||
model.load_device = config.compute_device
|
||||
|
||||
logging.info("[DisTorch] Block swap configuration applied successfully")
|
||||
try:
|
||||
# Create config
|
||||
config = BlockSwapConfig(
|
||||
virtual_vram_gb=virtual_vram_gb,
|
||||
swap_space_gb=swap_space_gb,
|
||||
swap_device=swap_device,
|
||||
compute_device=compute_device,
|
||||
use_non_blocking=use_non_blocking
|
||||
)
|
||||
|
||||
# Get the actual model (handle ModelPatcher)
|
||||
if hasattr(model, 'model'):
|
||||
actual_model = model.model
|
||||
distorch_logger.debug(f"Extracted actual model from ModelPatcher: {type(actual_model)}")
|
||||
else:
|
||||
actual_model = model
|
||||
distorch_logger.debug(f"Using model directly: {type(actual_model)}")
|
||||
|
||||
# Check if model has diffusion_model (common pattern)
|
||||
if hasattr(actual_model, 'diffusion_model'):
|
||||
target_model = actual_model.diffusion_model
|
||||
distorch_logger.debug(f"Found diffusion_model: {type(target_model)}")
|
||||
else:
|
||||
target_model = actual_model
|
||||
distorch_logger.debug(f"No diffusion_model found, using model as-is")
|
||||
|
||||
# Create block swap manager
|
||||
distorch_logger.debug("Creating BlockSwapManager...")
|
||||
manager = BlockSwapManager(target_model, config)
|
||||
|
||||
# Prepare model (move blocks to swap device)
|
||||
distorch_logger.debug("Preparing model...")
|
||||
manager.prepare()
|
||||
|
||||
# Store manager on model for later access
|
||||
model._block_swap_manager = manager
|
||||
distorch_logger.debug("Stored BlockSwapManager on model._block_swap_manager")
|
||||
|
||||
# Also set the load_device attribute if it exists
|
||||
if hasattr(model, 'load_device'):
|
||||
model.load_device = config.compute_device
|
||||
distorch_logger.debug(f"Set model.load_device to {config.compute_device}")
|
||||
|
||||
distorch_logger.info("Block swap configuration applied successfully")
|
||||
distorch_logger.info("="*60)
|
||||
|
||||
except Exception as e:
|
||||
distorch_logger.error(f"Error in apply_block_swap: {e}")
|
||||
distorch_logger.error(traceback.format_exc())
|
||||
raise
|
||||
|
||||
return (model,)
|
||||
|
||||
Reference in New Issue
Block a user