mhc-ddim 0.1.0__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- mhc_ddim-0.1.0/LICENSE +1 -0
- mhc_ddim-0.1.0/PKG-INFO +27 -0
- mhc_ddim-0.1.0/README.md +0 -0
- mhc_ddim-0.1.0/pyproject.toml +55 -0
- mhc_ddim-0.1.0/setup.cfg +4 -0
- mhc_ddim-0.1.0/src/mhc_ddim/__init__.py +0 -0
- mhc_ddim-0.1.0/src/mhc_ddim/cfg.py +48 -0
- mhc_ddim-0.1.0/src/mhc_ddim/dist_util.py +85 -0
- mhc_ddim-0.1.0/src/mhc_ddim/fp16_util.py +69 -0
- mhc_ddim-0.1.0/src/mhc_ddim/gaussian_diffusion.py +1076 -0
- mhc_ddim-0.1.0/src/mhc_ddim/logger.py +494 -0
- mhc_ddim-0.1.0/src/mhc_ddim/losses.py +54 -0
- mhc_ddim-0.1.0/src/mhc_ddim/nn.py +326 -0
- mhc_ddim-0.1.0/src/mhc_ddim/resample.py +182 -0
- mhc_ddim-0.1.0/src/mhc_ddim/respace.py +122 -0
- mhc_ddim-0.1.0/src/mhc_ddim/script_util.py +305 -0
- mhc_ddim-0.1.0/src/mhc_ddim/train.py +356 -0
- mhc_ddim-0.1.0/src/mhc_ddim/unet.py +656 -0
- mhc_ddim-0.1.0/src/mhc_ddim.egg-info/PKG-INFO +27 -0
- mhc_ddim-0.1.0/src/mhc_ddim.egg-info/SOURCES.txt +21 -0
- mhc_ddim-0.1.0/src/mhc_ddim.egg-info/dependency_links.txt +1 -0
- mhc_ddim-0.1.0/src/mhc_ddim.egg-info/requires.txt +4 -0
- mhc_ddim-0.1.0/src/mhc_ddim.egg-info/top_level.txt +1 -0
mhc_ddim-0.1.0/LICENSE
ADDED
|
@@ -0,0 +1 @@
|
|
|
1
|
+
Copyright (c) 2026 Yasin Ansari
|
mhc_ddim-0.1.0/PKG-INFO
ADDED
|
@@ -0,0 +1,27 @@
|
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
|
+
Name: mhc-ddim
|
|
3
|
+
Version: 0.1.0
|
|
4
|
+
Summary: DDIM and DDPM implementations with Manifold-Constrained Hyper-Connections (mHC)
|
|
5
|
+
Author-email: Yasin Ansari <yasinansari.work@gmail.com>
|
|
6
|
+
License: MIT
|
|
7
|
+
Project-URL: Homepage, https://github.com/TryingManV2/mhc-ddim
|
|
8
|
+
Project-URL: Repository, https://github.com/TryingManV2/mhc-ddim
|
|
9
|
+
Project-URL: Issues, https://github.com/TryingManV2/mhc-ddim/issues
|
|
10
|
+
Keywords: diffusion,ddpm,ddim,mHC,manifold-constrained hyper-connections,pytorch,mnist,deep-learning,generative-models
|
|
11
|
+
Classifier: Development Status :: 3 - Alpha
|
|
12
|
+
Classifier: Intended Audience :: Science/Research
|
|
13
|
+
Classifier: License :: OSI Approved :: MIT License
|
|
14
|
+
Classifier: Programming Language :: Python :: 3
|
|
15
|
+
Classifier: Programming Language :: Python :: 3.9
|
|
16
|
+
Classifier: Programming Language :: Python :: 3.10
|
|
17
|
+
Classifier: Programming Language :: Python :: 3.11
|
|
18
|
+
Classifier: Programming Language :: Python :: 3.12
|
|
19
|
+
Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
|
|
20
|
+
Requires-Python: >=3.9
|
|
21
|
+
Description-Content-Type: text/markdown
|
|
22
|
+
License-File: LICENSE
|
|
23
|
+
Requires-Dist: torch
|
|
24
|
+
Requires-Dist: numpy
|
|
25
|
+
Requires-Dist: matplotlib
|
|
26
|
+
Requires-Dist: tqdm
|
|
27
|
+
Dynamic: license-file
|
mhc_ddim-0.1.0/README.md
ADDED
|
File without changes
|
|
@@ -0,0 +1,55 @@
|
|
|
1
|
+
[build-system]
|
|
2
|
+
requires = ["setuptools>=61"]
|
|
3
|
+
build-backend = "setuptools.build_meta"
|
|
4
|
+
|
|
5
|
+
[project]
|
|
6
|
+
name = "mhc-ddim"
|
|
7
|
+
version = "0.1.0"
|
|
8
|
+
description = "DDIM and DDPM implementations with Manifold-Constrained Hyper-Connections (mHC)"
|
|
9
|
+
readme = "README.md"
|
|
10
|
+
requires-python = ">=3.9"
|
|
11
|
+
|
|
12
|
+
authors = [
|
|
13
|
+
{ name = "Yasin Ansari", email = "yasinansari.work@gmail.com" }
|
|
14
|
+
]
|
|
15
|
+
|
|
16
|
+
license = { text = "MIT" }
|
|
17
|
+
|
|
18
|
+
keywords = [
|
|
19
|
+
"diffusion",
|
|
20
|
+
"ddpm",
|
|
21
|
+
"ddim",
|
|
22
|
+
"mHC",
|
|
23
|
+
"manifold-constrained hyper-connections",
|
|
24
|
+
"pytorch",
|
|
25
|
+
"mnist",
|
|
26
|
+
"deep-learning",
|
|
27
|
+
"generative-models"
|
|
28
|
+
]
|
|
29
|
+
|
|
30
|
+
classifiers = [
|
|
31
|
+
"Development Status :: 3 - Alpha",
|
|
32
|
+
"Intended Audience :: Science/Research",
|
|
33
|
+
"License :: OSI Approved :: MIT License",
|
|
34
|
+
"Programming Language :: Python :: 3",
|
|
35
|
+
"Programming Language :: Python :: 3.9",
|
|
36
|
+
"Programming Language :: Python :: 3.10",
|
|
37
|
+
"Programming Language :: Python :: 3.11",
|
|
38
|
+
"Programming Language :: Python :: 3.12",
|
|
39
|
+
"Topic :: Scientific/Engineering :: Artificial Intelligence",
|
|
40
|
+
]
|
|
41
|
+
|
|
42
|
+
dependencies = [
|
|
43
|
+
"torch",
|
|
44
|
+
"numpy",
|
|
45
|
+
"matplotlib",
|
|
46
|
+
"tqdm",
|
|
47
|
+
]
|
|
48
|
+
|
|
49
|
+
[project.urls]
|
|
50
|
+
Homepage = "https://github.com/TryingManV2/mhc-ddim"
|
|
51
|
+
Repository = "https://github.com/TryingManV2/mhc-ddim"
|
|
52
|
+
Issues = "https://github.com/TryingManV2/mhc-ddim/issues"
|
|
53
|
+
|
|
54
|
+
[tool.setuptools.packages.find]
|
|
55
|
+
where = ["src"]
|
mhc_ddim-0.1.0/setup.cfg
ADDED
|
File without changes
|
|
@@ -0,0 +1,48 @@
|
|
|
1
|
+
import torch
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
class CFGModelWrapper:
|
|
5
|
+
"""
|
|
6
|
+
Wraps a conditional UNet to implement classifier-free guidance at sampling.
|
|
7
|
+
|
|
8
|
+
Runs the model twice per step — once with the given class labels, once
|
|
9
|
+
with the null class — then combines:
|
|
10
|
+
|
|
11
|
+
eps = eps_uncond + guidance_scale * (eps_cond - eps_uncond)
|
|
12
|
+
|
|
13
|
+
Set guidance_scale = 1.0 to recover plain conditional sampling.
|
|
14
|
+
Set guidance_scale = 0.0 to recover plain unconditional sampling.
|
|
15
|
+
|
|
16
|
+
Only valid for ModelMeanType.EPSILON. For START_X or PREVIOUS_X you
|
|
17
|
+
would need the analogous combination on those quantities.
|
|
18
|
+
"""
|
|
19
|
+
def __init__(self, model, null_class, guidance_scale=3.0):
|
|
20
|
+
self.model = model
|
|
21
|
+
self.null_class = null_class
|
|
22
|
+
self.guidance_scale = guidance_scale
|
|
23
|
+
|
|
24
|
+
def __call__(self, x, t, **kwargs):
|
|
25
|
+
y = kwargs.get("y", None)
|
|
26
|
+
if y is None:
|
|
27
|
+
# No label provided: behave like the plain unconditional model.
|
|
28
|
+
return self.model(x, t, **kwargs)
|
|
29
|
+
|
|
30
|
+
null_y = torch.full_like(y, self.null_class)
|
|
31
|
+
|
|
32
|
+
kwargs_cond = {**kwargs, "y": y}
|
|
33
|
+
kwargs_uncond = {**kwargs, "y": null_y}
|
|
34
|
+
|
|
35
|
+
out_cond = self.model(x, t, **kwargs_cond)
|
|
36
|
+
out_uncond = self.model(x, t, **kwargs_uncond)
|
|
37
|
+
|
|
38
|
+
# If the model outputs variance too (learn_sigma=True), the second
|
|
39
|
+
# half of the channels must NOT be guided. Split and re-combine.
|
|
40
|
+
if out_cond.shape[1] == 2 * x.shape[1]:
|
|
41
|
+
eps_c, var_c = out_cond.chunk(2, dim=1)
|
|
42
|
+
eps_u, var_u = out_uncond.chunk(2, dim=1)
|
|
43
|
+
eps = eps_u + self.guidance_scale * (eps_c - eps_u)
|
|
44
|
+
# Use conditional variance; using unconditional is also fine.
|
|
45
|
+
return torch.cat([eps, var_c], dim=1)
|
|
46
|
+
|
|
47
|
+
# Standard case: pure epsilon prediction.
|
|
48
|
+
return out_uncond + self.guidance_scale * (out_cond - out_uncond)
|
|
@@ -0,0 +1,85 @@
|
|
|
1
|
+
import io
|
|
2
|
+
import os
|
|
3
|
+
import socket
|
|
4
|
+
|
|
5
|
+
import blobfile as bf # type: ignore
|
|
6
|
+
from mpi4py import MPI # type: ignore
|
|
7
|
+
import torch
|
|
8
|
+
import torch.distributed as dist
|
|
9
|
+
|
|
10
|
+
# Change this to reflect your cluster layout.
|
|
11
|
+
# The GPU for a given rank is (rank % GPUS_PER_NODE).
|
|
12
|
+
GPUS_PER_NODE = 8
|
|
13
|
+
|
|
14
|
+
SETUP_RETRY_COUNT = 3
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def setup_dist():
|
|
18
|
+
"""
|
|
19
|
+
Setup a distributed process group.
|
|
20
|
+
"""
|
|
21
|
+
if dist.is_initialized():
|
|
22
|
+
return
|
|
23
|
+
|
|
24
|
+
comm = MPI.COMM_WORLD
|
|
25
|
+
backend = "gloo" if not torch.cuda.is_available() else "nccl"
|
|
26
|
+
|
|
27
|
+
if backend == "gloo":
|
|
28
|
+
hostname = "localhost"
|
|
29
|
+
else:
|
|
30
|
+
hostname = socket.gethostbyname(socket.getfqdn())
|
|
31
|
+
os.environ["MASTER_ADDR"] = comm.bcast(hostname, root=0)
|
|
32
|
+
os.environ["RANK"] = str(comm.rank)
|
|
33
|
+
os.environ["WORLD_SIZE"] = str(comm.size)
|
|
34
|
+
|
|
35
|
+
port = comm.bcast(_find_free_port(), root=0)
|
|
36
|
+
os.environ["MASTER_PORT"] = str(port)
|
|
37
|
+
dist.init_process_group(backend=backend, init_method="env://")
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def dev():
|
|
41
|
+
"""
|
|
42
|
+
Get the device to use for torch.distributed.
|
|
43
|
+
"""
|
|
44
|
+
if torch.cuda.is_available():
|
|
45
|
+
return torch.device(f"cuda:{MPI.COMM_WORLD.Get_rank() % GPUS_PER_NODE}")
|
|
46
|
+
return torch.device("cpu")
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def load_state_dict(path, **kwargs):
|
|
50
|
+
"""
|
|
51
|
+
Load a PyTorch file without redundant fetches across MPI ranks.
|
|
52
|
+
"""
|
|
53
|
+
if MPI.COMM_WORLD.Get_rank() == 0:
|
|
54
|
+
with bf.BlobFile(path, "rb") as f:
|
|
55
|
+
data = f.read()
|
|
56
|
+
else:
|
|
57
|
+
data = None
|
|
58
|
+
data = MPI.COMM_WORLD.bcast(data)
|
|
59
|
+
|
|
60
|
+
# PyTorch >= 2.6 defaults `weights_only` to True, which rejects the
|
|
61
|
+
# Python objects stored in optimizer checkpoints. We load our own
|
|
62
|
+
# checkpoints, so explicitly fall back to the permissive unpickler
|
|
63
|
+
# unless the caller already chose otherwise.
|
|
64
|
+
kwargs.setdefault("weights_only", False)
|
|
65
|
+
|
|
66
|
+
return torch.load(io.BytesIO(data), **kwargs)
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
def sync_params(params):
|
|
70
|
+
"""
|
|
71
|
+
Synchronize a sequence of Tensors across ranks from rank 0.
|
|
72
|
+
"""
|
|
73
|
+
for p in params:
|
|
74
|
+
with torch.no_grad():
|
|
75
|
+
dist.broadcast(p, 0)
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
def _find_free_port():
|
|
79
|
+
s = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
|
|
80
|
+
try:
|
|
81
|
+
s.bind(("", 0))
|
|
82
|
+
s.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
|
83
|
+
return s.getsockname()[1]
|
|
84
|
+
finally:
|
|
85
|
+
s.close()
|
|
@@ -0,0 +1,69 @@
|
|
|
1
|
+
import torch.nn as nn
|
|
2
|
+
from torch._utils import _flatten_dense_tensors, _unflatten_dense_tensors
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
def convert_module_to_f16(l):
|
|
6
|
+
if isinstance(l, (nn.Conv1d, nn.Conv2d, nn.Conv3d)):
|
|
7
|
+
assert l.weight is not None
|
|
8
|
+
l.weight.data = l.weight.data.half()
|
|
9
|
+
if l.bias is not None:
|
|
10
|
+
l.bias.data = l.bias.data.half()
|
|
11
|
+
|
|
12
|
+
def convert_module_to_f32(l):
|
|
13
|
+
if isinstance(l, (nn.Conv1d, nn.Conv2d, nn.Conv3d)):
|
|
14
|
+
assert l.weight is not None
|
|
15
|
+
l.weight.data = l.weight.data.float()
|
|
16
|
+
if l.bias is not None:
|
|
17
|
+
l.bias.data = l.bias.data.float()
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def make_master_params(model_params):
|
|
21
|
+
"""
|
|
22
|
+
Copy model parameters into a (differently-shaped) list of full-precision
|
|
23
|
+
parameters.
|
|
24
|
+
"""
|
|
25
|
+
master_params = _flatten_dense_tensors(
|
|
26
|
+
[param.detach().float() for param in model_params]
|
|
27
|
+
)
|
|
28
|
+
master_params = nn.Parameter(master_params)
|
|
29
|
+
master_params.requires_grad = True
|
|
30
|
+
return [master_params]
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def model_grads_to_master_grads(model_params, master_params):
|
|
34
|
+
"""
|
|
35
|
+
Copy the gradients from the model parameters into the master parameters
|
|
36
|
+
from make_master_params().
|
|
37
|
+
"""
|
|
38
|
+
master_params[0].grad = _flatten_dense_tensors(
|
|
39
|
+
[param.grad.data.detach().float() for param in model_params]
|
|
40
|
+
)
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
def master_params_to_model_params(model_params, master_params):
|
|
44
|
+
"""
|
|
45
|
+
Copy the master parameter data back into the model parameters.
|
|
46
|
+
"""
|
|
47
|
+
# Without copying to a list, if a generator is passed, this will
|
|
48
|
+
# silently not copy any parameters.
|
|
49
|
+
model_params = list(model_params)
|
|
50
|
+
|
|
51
|
+
for param, master_param in zip(
|
|
52
|
+
model_params, unflatten_master_params(model_params, master_params)
|
|
53
|
+
):
|
|
54
|
+
param.detach().copy_(master_param)
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
def unflatten_master_params(model_params, master_params):
|
|
58
|
+
"""
|
|
59
|
+
Unflatten the master parameters to look like model_params.
|
|
60
|
+
"""
|
|
61
|
+
return _unflatten_dense_tensors(master_params[0].detach(), tuple(tensor for tensor in model_params))
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
def zero_grad(model_params):
|
|
65
|
+
for param in model_params:
|
|
66
|
+
# Taken from https://pytorch.org/docs/stable/_modules/torch/optim/optimizer.html#Optimizer.add_param_group
|
|
67
|
+
if param.grad is not None:
|
|
68
|
+
param.grad.detach_()
|
|
69
|
+
param.grad.zero_()
|