diff --git a/tools/modules/diffusions/schedules.py b/tools/modules/diffusions/schedules.py index 48b9236..0471a64 100644 --- a/tools/modules/diffusions/schedules.py +++ b/tools/modules/diffusions/schedules.py @@ -46,7 +46,7 @@ def sigma_schedule(schedule='cosine', def linear_schedule(num_timesteps, init_beta, last_beta, **kwargs): scale = 1000.0 / num_timesteps init_beta = init_beta or scale * 0.0001 - ast_beta = last_beta or scale * 0.02 + last_beta = last_beta or scale * 0.02 return torch.linspace(init_beta, last_beta, num_timesteps, dtype=torch.float64) def logsnr_cosine_interp_schedule( diff --git a/tools/modules/unet/util.py b/tools/modules/unet/util.py index 2903715..d6ffe27 100644 --- a/tools/modules/unet/util.py +++ b/tools/modules/unet/util.py @@ -348,7 +348,7 @@ def __init__(self, in_channels, n_heads, d_head, stride=1, padding=0)) else: - self.proj_out = zero_module(nn.Linear(in_channels, inner_dim)) + self.proj_out = zero_module(nn.Linear(inner_dim, in_channels)) self.use_linear = use_linear def forward(self, x, context=None): diff --git a/utils/distributed.py b/utils/distributed.py index cba28ba..2d4872f 100644 --- a/utils/distributed.py +++ b/utils/distributed.py @@ -128,7 +128,7 @@ def reduce_dict(input_dict, group=None, reduction='mean', **kwargs): # ensure that the orders of keys are consistent across processes if isinstance(input_dict, OrderedDict): - keys = list(input_dict.keys) + keys = list(input_dict.keys()) else: keys = sorted(input_dict.keys()) vals = [input_dict[key] for key in keys] @@ -337,9 +337,9 @@ def symbolic(graph, input): return _split(input) @staticmethod - def symbolic(ctx, input): + def forward(ctx, input): return _split(input) - + @staticmethod def backward(ctx, grad_output): return _all_gather(grad_output)