update readme&OOM
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user