update readme&OOM

This commit is contained in:
靖渊
2024-04-22 16:19:03 +08:00
parent 4f4e516164
commit 8ba0b4c673
13 changed files with 43 additions and 34 deletions
@@ -5,10 +5,10 @@ import os.path
import random
from collections import OrderedDict
from PIL.Image import Image
import torch
import torch.nn.functional as F
from PIL.Image import Image
from scepter.modules.model.network.diffusion.diffusion import GaussianDiffusion
from scepter.modules.model.network.diffusion.schedules import noise_schedule
from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, MODELS,
@@ -274,7 +274,7 @@ class DiffusionInference():
module = self.load(module)
self.loaded_model[name] = module
return module
elif module['device'] == 'cpu':
elif module['device'] == 'cpu' or module['device'] == "offline":
module = self.load(module)
return module
else:
@@ -276,7 +276,7 @@ class LargenInference():
module = self.load(module)
self.loaded_model[name] = module
return module
elif module['device'] == 'cpu':
elif module['device'] == 'cpu' or module['device'] == 'offline':
module = self.load(module)
return module
else:
@@ -498,7 +498,7 @@ class LargenInference():
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=True)
skip_loaded=False)
# cond stage
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
@@ -513,7 +513,7 @@ class LargenInference():
self.dynamic_unload(self.cond_stage_model,
'cond_stage_model',
skip_loaded=True)
skip_loaded=False)
# get noise
seed = kwargs.pop('seed', -1)
@@ -582,7 +582,7 @@ class LargenInference():
x_samples = self.decode_first_stage(latent).float()
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=True)
skip_loaded=False)
images = torch.clamp((x_samples + 1.0) / 2.0, min=0.0, max=1.0)
if base_image is not None:
stitch_images = []
@@ -5,11 +5,11 @@ import os.path
import random
from collections import OrderedDict
from PIL.Image import Image
import torch
import torch.nn.functional as F
import torchvision.transforms.functional as TF
from PIL.Image import Image
from scepter.modules.model.network.diffusion.diffusion import GaussianDiffusion
from scepter.modules.model.network.diffusion.schedules import noise_schedule
from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, MODELS,
@@ -275,7 +275,7 @@ class StyleboothInference():
module = self.load(module)
self.loaded_model[name] = module
return module
elif module['device'] == 'cpu':
elif module['device'] == 'cpu' or module['device'] == "offline":
module = self.load(module)
return module
else:
+12 -4
View File
@@ -49,6 +49,12 @@ class CheckpointHook(Hook):
'',
'description':
'If save the best model, which order should be sorted, +/-!'
},
'DISABLE_SNAPSHOT': {
'value':
False,
'description':
'Skip to save snapshot checkpoint.'
}
}]
@@ -62,6 +68,7 @@ class CheckpointHook(Hook):
self.save_best_by = cfg.get('SAVE_BEST_BY', '')
self.push_to_hub = cfg.get('PUSH_TO_HUB', False)
self.hub_model_id = cfg.get('HUB_MODEL_ID', None)
self.disable_save_snapshot = cfg.get('DISABLE_SNAPSHOT', False)
self.last_ckpt = None
if self.save_best and not self.save_best_by:
warnings.warn(
@@ -110,10 +117,11 @@ class CheckpointHook(Hook):
solver.work_dir,
'checkpoints/{}-{}.pth'.format(self.save_name_prefix,
solver.total_iter + 1))
with FS.put_to(save_path) as local_path:
with open(local_path, 'wb') as f:
checkpoint = solver.save_checkpoint()
torch.save(checkpoint, f)
if not self.disable_save_snapshot:
with FS.put_to(save_path) as local_path:
with open(local_path, 'wb') as f:
checkpoint = solver.save_checkpoint()
torch.save(checkpoint, f)
from swift import SwiftModel
if isinstance(solver.model, SwiftModel):