Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
Vladimir Mandic
2024-09-28 09:02:51 -04:00
parent d51a695d70
commit c6e4aee7a4
4 changed files with 7 additions and 7 deletions
+1 -1
View File
@@ -112,7 +112,7 @@ def setup_logging():
}))
logging.basicConfig(level=logging.ERROR, format='%(asctime)s | %(name)s | %(levelname)s | %(module)s | %(message)s', handlers=[logging.NullHandler()]) # redirect default logger to null
pretty_install(console=console)
traceback_install(console=console, extra_lines=1, max_frames=10, width=console.width, word_wrap=False, indent_guides=False, suppress=[])
traceback_install(console=console, extra_lines=1, max_frames=16, width=console.width, word_wrap=False, indent_guides=False, suppress=[])
while log.hasHandlers() and len(log.handlers) > 0:
log.removeHandler(log.handlers[0])
+3 -3
View File
@@ -17,7 +17,7 @@ console = Console(log_time=True, tab_size=4, log_time_format='%H:%M:%S-%f', soft
}))
pretty_install(console=console)
traceback_install(console=console, extra_lines=1, width=console.width, word_wrap=False, indent_guides=False)
traceback_install(console=console, extra_lines=1, width=console.width, word_wrap=False, indent_guides=False, max_frames=16)
already_displayed = {}
@@ -38,7 +38,7 @@ def print_error_explanation(message):
def display(e: Exception, task, suppress=[]):
log.error(f"{task or 'error'}: {type(e).__name__}")
console.print_exception(show_locals=False, max_frames=10, extra_lines=1, suppress=suppress, theme="ansi_dark", word_wrap=False, width=console.width)
console.print_exception(show_locals=False, max_frames=16, extra_lines=1, suppress=suppress, theme="ansi_dark", word_wrap=False, width=console.width)
def display_once(e: Exception, task):
@@ -56,7 +56,7 @@ def run(code, task):
def exception(suppress=[]):
console.print_exception(show_locals=False, max_frames=10, extra_lines=2, suppress=suppress, theme="ansi_dark", word_wrap=False, width=min([console.width, 200]))
console.print_exception(show_locals=False, max_frames=16, extra_lines=2, suppress=suppress, theme="ansi_dark", word_wrap=False, width=min([console.width, 200]))
def profile(profiler, msg: str):
+1 -1
View File
@@ -78,7 +78,7 @@ def create_sampler(name, model):
shared.log.warning(f'AlphaVLLM-Lumina: sampler="{name}" unsupported')
return None
if not hasattr(model, 'scheduler_config'):
model.scheduler_config = sampler.sampler.config.copy()
model.scheduler_config = sampler.sampler.config.copy() if hasattr(sampler.sampler, 'config') else {}
model.scheduler = sampler.sampler
if hasattr(model, "prior_pipe") and hasattr(model.prior_pipe, "scheduler"):
model.prior_pipe.scheduler = sampler.sampler
+2 -2
View File
@@ -406,8 +406,8 @@ class VDMScheduler(SchedulerMixin, ConfigMixin):
sqrt_alpha_prod = torch.sqrt(torch.sigmoid(gamma))
sqrt_one_minus_alpha_prod = torch.sqrt(torch.sigmoid(-gamma)) # sqrt(sigma)
noisy_samples = sqrt_alpha_prod * original_samples + sqrt_one_minus_alpha_prod * noise
return noisy_samples
noisy_samples = original_samples * sqrt_alpha_prod + noise * sqrt_one_minus_alpha_prod
return noisy_samples.to(original_samples.dtype)
def get_velocity(self, sample: torch.Tensor, noise: torch.Tensor, timesteps: torch.Tensor) -> torch.Tensor:
gamma = self.log_snr(timesteps).to(sample.device)