OSCR

RAVEN: Robust, generalizable, multi-resolution structural MRI upsampling using autoencoders.

Code ↔ Paper

9 matches between paragraphs of the paper and lines of its authors' code, computed by the harvester (lexical-v1). Click a colored paragraph or line to see its counterpart.

The 9 matches
  1. [1] § Methods › Architecture › Encoder ↔ inference/raven_v1_3/runtime/taming/modules/diffusionmodules/model.py, lines 2529–2633 · score 0.90 · MaxPool3d, residual connection, CrossResolutionAttention3D, downsampled version, feature map, Conv3d
  2. [2] § Methods › Architecture › Discriminator ↔ inference/raven_v1_3/runtime/taming/modules/discriminator/model.py, lines 70–140 · score 0.65 · BatchNorm3d, LeakyReLU, Conv3d, kernel, layer, Discriminator
  3. [3] § Methods › Architecture › Discriminator ↔ inference/raven_v1_3/runtime/taming/modules/discriminator/model.py, lines 70–140 · score 0.64 · BatchNorm3d, LeakyReLU, Conv3d, filters, kernel, Discriminator
  4. [4] § Methods › Architecture › Encoder ↔ inference/raven_v1_3/runtime/taming/models/mri_autoencoders.py, lines 10–65 · score 0.63 · GroupNorm, SiLU, Conv3d, stacked, tensor, layers
  5. [5] § Methods › Architecture › Encoder ↔ inference/raven_v1_3/runtime/taming/modules/diffusionmodules/model.py, lines 1888–1975 · score 0.59 · Downsample3D, SiLU, Conv3d, blocks, kernel, padding
  6. [6] § Methods › Architecture › Encoder ↔ inference/raven_v1_3/runtime/taming/modules/diffusionmodules/model2.py, lines 1892–1979 · score 0.59 · Downsample3D, SiLU, Conv3d, blocks, kernel, padding
  7. [7] § Methods › Training procedure › Generation of LR-HR image pair ↔ inference/raven_v1_3/runtime/taming/models/autoencoders.py, lines 181–215 · score 0.58 · gaussian blurring, low resolution, BS, degraded, trilinear, tensor
  8. [8] § Methods › Architecture › Encoder ↔ inference/raven_v1_3/runtime/taming/models/mri_autoencoders.py, lines 10–65 · score 0.57 · GroupNorm, SiLU, Conv3d, kernel, stride
  9. [9] § Methods › Training procedure › Generation of LR-HR image pair ↔ inference/raven_v1_3/runtime/taming/modules/losses/contperceptual.py, lines 2109–2256 · score 0.52 · KL Divergence, Reconstruction loss, Fake, GAN, discriminator, Training

Paper

Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC

The paper is loaded when this pane is shown.

The authors' code

Python · 2,743 lines · 113 KB · MIT · 2 matches

  1. # pytorch_diffusion + derived encoder decoder
  2. import math
  3. import torch
  4. import torch.nn as nn
  5. import numpy as np
  6. import torch.nn.functional as F
  7. import torch, warnings, functools, os
  8. def term_color(text, c): # quick & dirty ANSI colours
  9. codes = dict(red=31, green=32, yellow=33, cyan=36)
  10. return f"\033[{codes[c]}m{text}\033[0m"
  11. class FlashCompatMixin:
  12. """
  13. Adds `_flash_ok` property and `_explain_flash()` helper.
  14. Call `self._explain_flash()` once per module (e.g. in `__init__`).
  15. """
  16. def _flash_capable(self, head_dim, dtype):
  17. cc_ok = torch.cuda.get_device_capability()[0] >= 8
  18. dtype_ok= dtype in (torch.float16, torch.bfloat16)
  19. dim_ok = (head_dim <= 128) and (head_dim % 8 == 0)
  20. return cc_ok and dtype_ok and dim_ok
  21. def _explain_flash(self, head_dim, num_heads, dtype):
  22. if not torch.cuda.is_available():
  23. print(term_color("→ FlashAttention not available (CPU run).", "yellow"))
  24. return False
  25. flash = self._flash_capable(head_dim, dtype)
  26. kernel = torch.backends.cuda.preferred_linalg_library
  27. if flash:
  28. print(term_color(
  29. f"✓ Flash-SDP kernel will be used "
  30. f"(h={num_heads}, d={head_dim}, {dtype}, {kernel})", "green"))
  31. else:
  32. why = []
  33. cc = torch.cuda.get_device_capability()[0]
  34. if cc < 8: why.append(f"SM{cc*10} GPU")
  35. if dtype not in (torch.float16, torch.bfloat16): why.append(f"dtype={dtype}")
  36. if head_dim > 128: why.append(f"head_dim={head_dim}>128")
  37. if head_dim % 8: why.append(f"head_dim%8={head_dim%8}")
  38. msg = " / ".join(why)
  39. print(term_color(f"→ Flash disabled, falling back to efficient/math ({msg})",
  40. "red"))
  41. return flash
  42. def get_timestep_embedding(timesteps, embedding_dim):
  43. """
  44. This matches the implementation in Denoising Diffusion Probabilistic Models:
  45. From Fairseq.
  46. Build sinusoidal embeddings.
  47. This matches the implementation in tensor2tensor, but differs slightly
  48. from the description in Section 3.5 of "Attention Is All You Need".
  49. """
  50. assert len(timesteps.shape) == 1
  51. half_dim = embedding_dim // 2
  52. emb = math.log(10000) / (half_dim - 1)
  53. emb = torch.exp(torch.arange(half_dim, dtype=torch.float32) * -emb)
  54. emb = emb.to(device=timesteps.device)
  55. emb = timesteps.float()[:, None] * emb[None, :]
  56. emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=1)
  57. if embedding_dim % 2 == 1: # zero pad
  58. emb = torch.nn.functional.pad(emb, (0,1,0,0))
  59. return emb
  60. def nonlinearity(x):
  61. # swish
  62. return x*torch.sigmoid(x)
  63. def Normalize(in_channels):
  64. return torch.nn.GroupNorm(num_groups=8, num_channels=in_channels, eps=1e-6, affine=True)
  65. class Upsample(nn.Module):
  66. def __init__(self, in_channels, with_conv):
  67. super().__init__()
  68. self.with_conv = with_conv
  69. if self.with_conv:
  70. self.conv = torch.nn.Conv2d(in_channels,
  71. in_channels,
  72. kernel_size=3,
  73. stride=1,
  74. padding=1)
  75. def forward(self, x):
  76. x = torch.nn.functional.interpolate(x, scale_factor=2.0, mode="nearest")
  77. if self.with_conv:
  78. x = self.conv(x)
  79. return x
  80. class Downsample(nn.Module):
  81. def __init__(self, in_channels, with_conv):
  82. super().__init__()
  83. self.with_conv = with_conv
  84. if self.with_conv:
  85. # no asymmetric padding in torch conv, must do it ourselves
  86. self.conv = torch.nn.Conv2d(in_channels,
  87. in_channels,
  88. kernel_size=3,
  89. stride=2,
  90. padding=0)
  91. def forward(self, x):
  92. if self.with_conv:
  93. pad = (0,1,0,1)
  94. x = torch.nn.functional.pad(x, pad, mode="constant", value=0)
  95. x = self.conv(x)
  96. else:
  97. x = torch.nn.functional.avg_pool2d(x, kernel_size=2, stride=2)
  98. return x
  99. class Upsample3D(nn.Module):
  100. def __init__(self, in_channels, with_conv):
  101. """
  102. 3D version of your Upsample block.
  103. Assumes input shape [batch_size, in_channels, D, H, W].
  104. """
  105. super().__init__()
  106. self.with_conv = with_conv
  107. if self.with_conv:
  108. self.conv = nn.Conv3d(
  109. in_channels,
  110. in_channels,
  111. kernel_size=3,
  112. stride=1,
  113. padding=1
  114. )
  115. def forward(self, x):
  116. # For 3D data, use 3D interpolation
  117. # scale_factor=(2,2,2) will upsample D, H, W all by factor of 2
  118. x = F.interpolate(x, scale_factor=(2, 2, 2), mode="nearest")
  119. if self.with_conv:
  120. x = self.conv(x)
  121. return x
  122. class Downsample3D(nn.Module):
  123. def __init__(self, in_channels, with_conv):
  124. """
  125. 3D version of your Downsample block.
  126. Assumes input shape [batch_size, in_channels, D, H, W].
  127. """
  128. super().__init__()
  129. self.with_conv = with_conv
  130. if self.with_conv:
  131. # No asymmetric padding in torch Conv3d, must do it ourselves
  132. # kernel_size=3, stride=2, padding=0
  133. self.conv = nn.Conv3d(
  134. in_channels,
  135. in_channels,
  136. kernel_size=3,
  137. stride=2,
  138. padding=0
  139. )
  140. def forward(self, x):
  141. if self.with_conv:
  142. # 3D padding tuple: (pad_left, pad_right, pad_top, pad_bottom, pad_front, pad_back)
  143. # e.g. we add 1 unit of padding to right, bottom, and back if needed
  144. pad = (0, 1, 0, 1, 0, 1)
  145. x = F.pad(x, pad, mode="constant", value=0)
  146. x = self.conv(x)
  147. else:
  148. # If no conv, just average-pool with kernel_size=2
  149. x = F.avg_pool3d(x, kernel_size=2, stride=2)
  150. return x
  151. class ResnetBlock(nn.Module):
  152. def __init__(self, *, in_channels, out_channels=None, conv_shortcut=False,
  153. dropout, temb_channels=512):
  154. super().__init__()
  155. self.in_channels = in_channels
  156. out_channels = in_channels if out_channels is None else out_channels
  157. self.out_channels = out_channels
  158. self.use_conv_shortcut = conv_shortcut
  159. self.norm1 = Normalize(in_channels)
  160. self.conv1 = torch.nn.Conv2d(in_channels,
  161. out_channels,
  162. kernel_size=3,
  163. stride=1,
  164. padding=1)
  165. if temb_channels > 0:
  166. self.temb_proj = torch.nn.Linear(temb_channels,
  167. out_channels)
  168. self.norm2 = Normalize(out_channels)
  169. self.dropout = torch.nn.Dropout(dropout)
  170. self.conv2 = torch.nn.Conv2d(out_channels,
  171. out_channels,
  172. kernel_size=3,
  173. stride=1,
  174. padding=1)
  175. if self.in_channels != self.out_channels:
  176. if self.use_conv_shortcut:
  177. self.conv_shortcut = torch.nn.Conv2d(in_channels,
  178. out_channels,
  179. kernel_size=3,
  180. stride=1,
  181. padding=1)
  182. else:
  183. self.nin_shortcut = torch.nn.Conv2d(in_channels,
  184. out_channels,
  185. kernel_size=1,
  186. stride=1,
  187. padding=0)
  188. def forward(self, x, temb):
  189. h = x
  190. h = self.norm1(h)
  191. h = nonlinearity(h)
  192. h = self.conv1(h)
  193. if temb is not None:
  194. h = h + self.temb_proj(nonlinearity(temb))[:,:,None,None]
  195. h = self.norm2(h)
  196. h = nonlinearity(h)
  197. h = self.dropout(h)
  198. h = self.conv2(h)
  199. if self.in_channels != self.out_channels:
  200. if self.use_conv_shortcut:
  201. x = self.conv_shortcut(x)
  202. else:
  203. x = self.nin_shortcut(x)
  204. return x+h
  205. class ResnetBlock3D(nn.Module):
  206. def __init__(self, *, in_channels, out_channels=None, conv_shortcut=False,
  207. dropout, temb_channels=512):
  208. super().__init__()
  209. self.in_channels = in_channels
  210. out_channels = in_channels if out_channels is None else out_channels
  211. self.out_channels = out_channels
  212. self.use_conv_shortcut = conv_shortcut
  213. self.norm1 = Normalize(in_channels)
  214. self.conv1 = torch.nn.Conv3d(in_channels,
  215. out_channels,
  216. kernel_size=3,
  217. stride=1,
  218. padding=1)
  219. if temb_channels > 0:
  220. self.temb_proj = torch.nn.Linear(temb_channels,
  221. out_channels)
  222. self.norm2 = Normalize(out_channels)
  223. self.dropout = torch.nn.Dropout(dropout)
  224. self.conv2 = torch.nn.Conv3d(out_channels,
  225. out_channels,
  226. kernel_size=3,
  227. stride=1,
  228. padding=1)
  229. if self.in_channels != self.out_channels:
  230. if self.use_conv_shortcut:
  231. self.conv_shortcut = torch.nn.Conv3d(in_channels,
  232. out_channels,
  233. kernel_size=3,
  234. stride=1,
  235. padding=1)
  236. else:
  237. self.nin_shortcut = torch.nn.Conv3d(in_channels,
  238. out_channels,
  239. kernel_size=1,
  240. stride=1,
  241. padding=0)
  242. def forward(self, x, temb):
  243. h = x
  244. h = self.norm1(h)
  245. h = nonlinearity(h)
  246. h = self.conv1(h)
  247. if temb is not None:
  248. h = h + self.temb_proj(nonlinearity(temb))[:,:,None,None]
  249. h = self.norm2(h)
  250. h = nonlinearity(h)
  251. h = self.dropout(h)
  252. h = self.conv2(h)
  253. if self.in_channels != self.out_channels:
  254. if self.use_conv_shortcut:
  255. x = self.conv_shortcut(x)
  256. else:
  257. x = self.nin_shortcut(x)
  258. return x+h
  259. class AttnBlock(nn.Module):
  260. def __init__(self, in_channels):
  261. super().__init__()
  262. self.in_channels = in_channels
  263. self.norm = Normalize(in_channels)
  264. self.q = torch.nn.Conv2d(in_channels,
  265. in_channels,
  266. kernel_size=1,
  267. stride=1,
  268. padding=0)
  269. self.k = torch.nn.Conv2d(in_channels,
  270. in_channels,
  271. kernel_size=1,
  272. stride=1,
  273. padding=0)
  274. self.v = torch.nn.Conv2d(in_channels,
  275. in_channels,
  276. kernel_size=1,
  277. stride=1,
  278. padding=0)
  279. self.proj_out = torch.nn.Conv2d(in_channels,
  280. in_channels,
  281. kernel_size=1,
  282. stride=1,
  283. padding=0)
  284. def forward(self, x):
  285. h_ = x
  286. h_ = self.norm(h_)
  287. q = self.q(h_)
  288. k = self.k(h_)
  289. v = self.v(h_)
  290. # compute attention
  291. b,c,h,w = q.shape
  292. q = q.reshape(b,c,h*w)
  293. q = q.permute(0,2,1) # b,hw,c
  294. k = k.reshape(b,c,h*w) # b,c,hw
  295. w_ = torch.bmm(q,k) # b,hw,hw w[b,i,j]=sum_c q[b,i,c]k[b,c,j]
  296. w_ = w_ * (int(c)**(-0.5))
  297. w_ = torch.nn.functional.softmax(w_, dim=2)
  298. # attend to values
  299. v = v.reshape(b,c,h*w)
  300. w_ = w_.permute(0,2,1) # b,hw,hw (first hw of k, second of q)
  301. h_ = torch.bmm(v,w_) # b, c,hw (hw of q) h_[b,c,j] = sum_i v[b,c,i] w_[b,i,j]
  302. h_ = h_.reshape(b,c,h,w)
  303. h_ = self.proj_out(h_)
  304. return x+h_
  305. class AttnBlock3D(nn.Module):
  306. def __init__(self, in_channels):
  307. super().__init__()
  308. self.in_channels = in_channels
  309. self.norm = Normalize(in_channels)
  310. self.q = torch.nn.Conv3d(in_channels,
  311. in_channels,
  312. kernel_size=1,
  313. stride=1,
  314. padding=0)
  315. self.k = torch.nn.Conv3d(in_channels,
  316. in_channels,
  317. kernel_size=1,
  318. stride=1,
  319. padding=0)
  320. self.v = torch.nn.Conv3d(in_channels,
  321. in_channels,
  322. kernel_size=1,
  323. stride=1,
  324. padding=0)
  325. self.proj_out = torch.nn.Conv3d(in_channels,
  326. in_channels,
  327. kernel_size=1,
  328. stride=1,
  329. padding=0)
  330. def forward(self, x):
  331. h_ = x
  332. h_ = self.norm(h_)
  333. q = self.q(h_)
  334. k = self.k(h_)
  335. v = self.v(h_)
  336. # compute attention
  337. b,c,h,w,z = q.shape
  338. q = q.reshape(b,c,h*w*z)
  339. q = q.permute(0,2,1) # b,hwz,c
  340. k = k.reshape(b,c,h*w*z) # b,c,hwz
  341. w_ = torch.bmm(q,k) # b,hwz,hwz w[b,i,j,k]=sum_c q[b,i,c]k[b,c,j]
  342. w_ = w_ * (int(c)**(-0.5))
  343. w_ = torch.nn.functional.softmax(w_, dim=2)
  344. # attend to values
  345. v = v.reshape(b,c,h*w*z)
  346. w_ = w_.permute(0,2,1) # b,hwz,hwz (first hw of k, second of q)
  347. h_ = torch.bmm(v,w_) # b, c,hwz (hw of q) h_[b,c,j] = sum_i v[b,c,i] w_[b,i,j]
  348. h_ = h_.reshape(b,c,h,w,z)
  349. h_ = self.proj_out(h_)
  350. return x+h_
  351. class Model(nn.Module):
  352. def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
  353. attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels,
  354. resolution, use_timestep=True):
  355. super().__init__()
  356. self.ch = ch
  357. self.temb_ch = self.ch*4
  358. self.num_resolutions = len(ch_mult)
  359. self.num_res_blocks = num_res_blocks
  360. self.resolution = resolution
  361. self.in_channels = in_channels
  362. self.use_timestep = use_timestep
  363. if self.use_timestep:
  364. # timestep embedding
  365. self.temb = nn.Module()
  366. self.temb.dense = nn.ModuleList([
  367. torch.nn.Linear(self.ch,
  368. self.temb_ch),
  369. torch.nn.Linear(self.temb_ch,
  370. self.temb_ch),
  371. ])
  372. # downsampling
  373. self.conv_in = torch.nn.Conv2d(in_channels,
  374. self.ch,
  375. kernel_size=3,
  376. stride=1,
  377. padding=1)
  378. curr_res = resolution
  379. in_ch_mult = (1,)+tuple(ch_mult)
  380. self.down = nn.ModuleList()
  381. for i_level in range(self.num_resolutions):
  382. block = nn.ModuleList()
  383. attn = nn.ModuleList()
  384. block_in = ch*in_ch_mult[i_level]
  385. block_out = ch*ch_mult[i_level]
  386. for i_block in range(self.num_res_blocks):
  387. block.append(ResnetBlock(in_channels=block_in,
  388. out_channels=block_out,
  389. temb_channels=self.temb_ch,
  390. dropout=dropout))
  391. block_in = block_out
  392. if curr_res in attn_resolutions:
  393. attn.append(AttnBlock(block_in))
  394. down = nn.Module()
  395. down.block = block
  396. down.attn = attn
  397. if i_level != self.num_resolutions-1:
  398. down.downsample = Downsample(block_in, resamp_with_conv)
  399. curr_res = curr_res // 2
  400. self.down.append(down)
  401. # middle
  402. self.mid = nn.Module()
  403. self.mid.block_1 = ResnetBlock(in_channels=block_in,
  404. out_channels=block_in,
  405. temb_channels=self.temb_ch,
  406. dropout=dropout)
  407. self.mid.attn_1 = AttnBlock(block_in)
  408. self.mid.block_2 = ResnetBlock(in_channels=block_in,
  409. out_channels=block_in,
  410. temb_channels=self.temb_ch,
  411. dropout=dropout)
  412. # upsampling
  413. self.up = nn.ModuleList()
  414. for i_level in reversed(range(self.num_resolutions)):
  415. block = nn.ModuleList()
  416. attn = nn.ModuleList()
  417. block_out = ch*ch_mult[i_level]
  418. skip_in = ch*ch_mult[i_level]
  419. for i_block in range(self.num_res_blocks+1):
  420. if i_block == self.num_res_blocks:
  421. skip_in = ch*in_ch_mult[i_level]
  422. block.append(ResnetBlock(in_channels=block_in+skip_in,
  423. out_channels=block_out,
  424. temb_channels=self.temb_ch,
  425. dropout=dropout))
  426. block_in = block_out
  427. if curr_res in attn_resolutions:
  428. attn.append(AttnBlock(block_in))
  429. up = nn.Module()
  430. up.block = block
  431. up.attn = attn
  432. if i_level != 0:
  433. up.upsample = Upsample(block_in, resamp_with_conv)
  434. curr_res = curr_res * 2
  435. self.up.insert(0, up) # prepend to get consistent order
  436. # end
  437. self.norm_out = Normalize(block_in)
  438. self.conv_out = torch.nn.Conv2d(block_in,
  439. out_ch,
  440. kernel_size=3,
  441. stride=1,
  442. padding=1)
  443. def forward(self, x, t=None):
  444. #assert x.shape[2] == x.shape[3] == self.resolution
  445. if self.use_timestep:
  446. # timestep embedding
  447. assert t is not None
  448. temb = get_timestep_embedding(t, self.ch)
  449. temb = self.temb.dense[0](temb)
  450. temb = nonlinearity(temb)
  451. temb = self.temb.dense[1](temb)
  452. else:
  453. temb = None
  454. # downsampling
  455. hs = [self.conv_in(x)]
  456. for i_level in range(self.num_resolutions):
  457. for i_block in range(self.num_res_blocks):
  458. h = self.down[i_level].block[i_block](hs[-1], temb)
  459. if len(self.down[i_level].attn) > 0:
  460. h = self.down[i_level].attn[i_block](h)
  461. hs.append(h)
  462. if i_level != self.num_resolutions-1:
  463. hs.append(self.down[i_level].downsample(hs[-1]))
  464. # middle
  465. h = hs[-1]
  466. h = self.mid.block_1(h, temb)
  467. h = self.mid.attn_1(h)
  468. h = self.mid.block_2(h, temb)
  469. # upsampling
  470. for i_level in reversed(range(self.num_resolutions)):
  471. for i_block in range(self.num_res_blocks+1):
  472. h = self.up[i_level].block[i_block](
  473. torch.cat([h, hs.pop()], dim=1), temb)
  474. if len(self.up[i_level].attn) > 0:
  475. h = self.up[i_level].attn[i_block](h)
  476. if i_level != 0:
  477. h = self.up[i_level].upsample(h)
  478. # end
  479. h = self.norm_out(h)
  480. h = nonlinearity(h)
  481. h = self.conv_out(h)
  482. return h
  483. class Model3D(nn.Module):
  484. def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
  485. attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels,
  486. resolution, use_timestep=True):
  487. super().__init__()
  488. self.ch = ch
  489. self.temb_ch = self.ch*4
  490. self.num_resolutions = len(ch_mult)
  491. self.num_res_blocks = num_res_blocks
  492. self.resolution = resolution
  493. self.in_channels = in_channels
  494. self.use_timestep = use_timestep
  495. if self.use_timestep:
  496. # timestep embedding
  497. self.temb = nn.Module()
  498. self.temb.dense = nn.ModuleList([
  499. torch.nn.Linear(self.ch,
  500. self.temb_ch),
  501. torch.nn.Linear(self.temb_ch,
  502. self.temb_ch),
  503. ])
  504. # downsampling
  505. self.conv_in = torch.nn.Conv2d(in_channels,
  506. self.ch,
  507. kernel_size=3,
  508. stride=1,
  509. padding=1)
  510. curr_res = resolution
  511. in_ch_mult = (1,)+tuple(ch_mult)
  512. self.down = nn.ModuleList()
  513. for i_level in range(self.num_resolutions):
  514. block = nn.ModuleList()
  515. attn = nn.ModuleList()
  516. block_in = ch*in_ch_mult[i_level]
  517. block_out = ch*ch_mult[i_level]
  518. for i_block in range(self.num_res_blocks):
  519. block.append(ResnetBlock3D(in_channels=block_in,
  520. out_channels=block_out,
  521. temb_channels=self.temb_ch,
  522. dropout=dropout))
  523. block_in = block_out
  524. if curr_res in attn_resolutions:
  525. attn.append(AttnBlock3D(block_in))
  526. down = nn.Module()
  527. down.block = block
  528. down.attn = attn
  529. if i_level != self.num_resolutions-1:
  530. down.downsample = Downsample3D(block_in, resamp_with_conv)
  531. curr_res = curr_res // 2
  532. self.down.append(down)
  533. # middle
  534. self.mid = nn.Module()
  535. self.mid.block_1 = ResnetBlock3D(in_channels=block_in,
  536. out_channels=block_in,
  537. temb_channels=self.temb_ch,
  538. dropout=dropout)
  539. self.mid.attn_1 = AttnBlock3D(block_in)
  540. self.mid.block_2 = ResnetBlock3D(in_channels=block_in,
  541. out_channels=block_in,
  542. temb_channels=self.temb_ch,
  543. dropout=dropout)
  544. # upsampling
  545. self.up = nn.ModuleList()
  546. for i_level in reversed(range(self.num_resolutions)):
  547. block = nn.ModuleList()
  548. attn = nn.ModuleList()
  549. block_out = ch*ch_mult[i_level]
  550. skip_in = ch*ch_mult[i_level]
  551. for i_block in range(self.num_res_blocks+1):
  552. if i_block == self.num_res_blocks:
  553. skip_in = ch*in_ch_mult[i_level]
  554. block.append(ResnetBlock3D(in_channels=block_in+skip_in,
  555. out_channels=block_out,
  556. temb_channels=self.temb_ch,
  557. dropout=dropout))
  558. block_in = block_out
  559. if curr_res in attn_resolutions:
  560. attn.append(AttnBlock3D(block_in))
  561. up = nn.Module()
  562. up.block = block
  563. up.attn = attn
  564. if i_level != 0:
  565. up.upsample = Upsample3D(block_in, resamp_with_conv)
  566. curr_res = curr_res * 2
  567. self.up.insert(0, up) # prepend to get consistent order
  568. # end
  569. self.norm_out = Normalize(block_in)
  570. self.conv_out = torch.nn.Conv3d(block_in,
  571. out_ch,
  572. kernel_size=3,
  573. stride=1,
  574. padding=1)
  575. def forward(self, x, t=None):
  576. #assert x.shape[2] == x.shape[3] == self.resolution
  577. if self.use_timestep:
  578. # timestep embedding
  579. assert t is not None
  580. temb = get_timestep_embedding(t, self.ch)
  581. temb = self.temb.dense[0](temb)
  582. temb = nonlinearity(temb)
  583. temb = self.temb.dense[1](temb)
  584. else:
  585. temb = None
  586. # downsampling
  587. hs = [self.conv_in(x)]
  588. for i_level in range(self.num_resolutions):
  589. for i_block in range(self.num_res_blocks):
  590. h = self.down[i_level].block[i_block](hs[-1], temb)
  591. if len(self.down[i_level].attn) > 0:
  592. h = self.down[i_level].attn[i_block](h)
  593. hs.append(h)
  594. if i_level != self.num_resolutions-1:
  595. hs.append(self.down[i_level].downsample(hs[-1]))
  596. # middle
  597. h = hs[-1]
  598. h = self.mid.block_1(h, temb)
  599. h = self.mid.attn_1(h)
  600. h = self.mid.block_2(h, temb)
  601. # upsampling
  602. for i_level in reversed(range(self.num_resolutions)):
  603. for i_block in range(self.num_res_blocks+1):
  604. h = self.up[i_level].block[i_block](
  605. torch.cat([h, hs.pop()], dim=1), temb)
  606. if len(self.up[i_level].attn) > 0:
  607. h = self.up[i_level].attn[i_block](h)
  608. if i_level != 0:
  609. h = self.up[i_level].upsample(h)
  610. # end
  611. h = self.norm_out(h)
  612. h = nonlinearity(h)
  613. h = self.conv_out(h)
  614. return h
  615. def rescale_to_uniform_voxel_size(x0_0, orig_zoom, iz_h=1.0, iz_w=1.0):
  616. def ceil_next_even(f):
  617. if f == 0:
  618. return 32 # Minimum dimension
  619. else:
  620. return 32 * torch.ceil(f / 32.)
  621. batch_size, filtnum, h, w = x0_0.size()
  622. # Lists to store results
  623. onemm_slices = []
  624. inner_network_shape_perbs = []
  625. padding_info = [] # Store padding values
  626. # Loop through each batch member
  627. for bs in range(batch_size):
  628. # Get the original zoom factors for each dimension
  629. orig_vox_h, orig_vox_w, _ = orig_zoom[bs, :]
  630. # Compute rescaling factors inversely proportional to iz_h and iz_w
  631. factor_h = orig_vox_h / iz_h # Inverse relationship
  632. factor_w = orig_vox_w / iz_w # Inverse relationship
  633. # Compute the target size for the current batch member using ceil_next_even
  634. H_norm = int(ceil_next_even(h * factor_h).item()) # Convert tensor to scalar
  635. W_norm = int(ceil_next_even(w * factor_w).item()) # Convert tensor to scalar
  636. # Resize current slice to the new normalized size
  637. resized_slice = F.interpolate(x0_0[bs:bs+1], size=(H_norm, W_norm), mode='bicubic', align_corners=True)
  638. # Clamp values between 0 and 1
  639. resized_slice = torch.clamp(resized_slice, 0, 1)
  640. onemm_slices.append(resized_slice)
  641. inner_network_shape_perbs.append((H_norm, W_norm))
  642. # Compute the maximum normalized height and width to zero-pad slices
  643. H_norm_max = max([shape[0] for shape in inner_network_shape_perbs])
  644. W_norm_max = max([shape[1] for shape in inner_network_shape_perbs])
  645. # Create a tensor to store the padded slices
  646. x0_00 = torch.zeros((batch_size, filtnum, H_norm_max, W_norm_max), device=x0_0.device)
  647. # Zero-pad each resized slice to match the max dimensions and store padding info
  648. for bs in range(batch_size):
  649. resized_slice = onemm_slices[bs]
  650. _, _, H_norm, W_norm = resized_slice.size()
  651. # Compute padding values
  652. pad_h = (H_norm_max - H_norm) // 2
  653. pad_w = (W_norm_max - W_norm) // 2
  654. # Pad the slice
  655. x0_00[bs] = F.pad(resized_slice, (pad_w, pad_w, pad_h, pad_h), mode='constant', value=0)
  656. # Store padding information for this batch element
  657. padding_info.append((pad_h, pad_h, pad_w, pad_w))
  658. # Return the padded tensor, inner shape per batch, and the padding info
  659. return x0_00, inner_network_shape_perbs, padding_info
  660. def invert_rescale_to_uniform_voxel_size(x0_00, inner_network_shape_perbs, padding_info, output_shape):
  661. batch_size, filtnum, H_norm_max, W_norm_max = x0_00.size()
  662. # List to store the unpadded and resized slices
  663. unpadded_slices = []
  664. # Loop through each batch member to remove padding and resize
  665. for bs in range(batch_size):
  666. # Get the padding information and original inner shape for this batch member
  667. pad_h, _, pad_w, _ = padding_info[bs]
  668. H_norm, W_norm = inner_network_shape_perbs[bs]
  669. # Calculate the start and end indices for slicing, based on padding
  670. start_h = pad_h
  671. end_h = start_h + H_norm
  672. start_w = pad_w
  673. end_w = start_w + W_norm
  674. # Slice the tensor based on the computed indices
  675. unpadded_slice = x0_00[bs, :, start_h:end_h, start_w:end_w]
  676. # Resize back to the original dimensions using output_shape
  677. unpadded_resized_slice = F.interpolate(unpadded_slice.unsqueeze(0), size=output_shape, mode='bicubic', align_corners=True)
  678. # Store the final result
  679. unpadded_slices.append(unpadded_resized_slice)
  680. # Concatenate all unpadded and resized slices into a final tensor
  681. return torch.cat(unpadded_slices, dim=0)
  682. class EncoderVINN(nn.Module):
  683. def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
  684. attn_resolutions, dropout=0.0, resamp_with_conv=True,
  685. inner_res = 0.7,
  686. in_channels, z_channels, double_z=True, **ignore_kwargs):
  687. super().__init__()
  688. self.ch = ch
  689. self.temb_ch = 0
  690. self.num_resolutions = len(ch_mult)
  691. self.num_res_blocks = num_res_blocks
  692. self.in_channels = in_channels
  693. self.inner_res = inner_res
  694. # downsampling
  695. self.conv_in = torch.nn.Conv2d(in_channels,
  696. self.ch,
  697. kernel_size=3,
  698. stride=1,
  699. padding=1)
  700. self.conv_vinn = torch.nn.Conv2d(self.ch,
  701. self.ch,
  702. kernel_size=3,
  703. stride=1,
  704. padding=1)
  705. curr_res = inner_res
  706. in_ch_mult = (1,)+tuple(ch_mult)
  707. self.down = nn.ModuleList()
  708. for i_level in range(self.num_resolutions):
  709. block = nn.ModuleList()
  710. attn = nn.ModuleList()
  711. block_in = ch*in_ch_mult[i_level]
  712. block_out = ch*ch_mult[i_level]
  713. for i_block in range(self.num_res_blocks):
  714. block.append(ResnetBlock(in_channels=block_in,
  715. out_channels=block_out,
  716. temb_channels=self.temb_ch,
  717. dropout=dropout))
  718. block_in = block_out
  719. if curr_res in attn_resolutions: # WATCH OUT! ATTN RESOLUTIONS ARE NOW IN TERMS OF VOXEL SIZE OF FEATURE MAPS
  720. attn.append(AttnBlock(block_in))
  721. down = nn.Module()
  722. down.block = block
  723. down.attn = attn
  724. if i_level != self.num_resolutions-1:
  725. down.downsample = Downsample(block_in, resamp_with_conv)
  726. curr_res = curr_res * 2
  727. self.down.append(down)
  728. # middle
  729. self.mid = nn.Module()
  730. self.mid.block_1 = ResnetBlock(in_channels=block_in,
  731. out_channels=block_in,
  732. temb_channels=self.temb_ch,
  733. dropout=dropout)
  734. self.mid.attn_1 = AttnBlock(block_in)
  735. self.mid.block_2 = ResnetBlock(in_channels=block_in,
  736. out_channels=block_in,
  737. temb_channels=self.temb_ch,
  738. dropout=dropout)
  739. # end
  740. self.norm_out = Normalize(block_in)
  741. self.conv_out = torch.nn.Conv2d(block_in,
  742. 2*z_channels if double_z else z_channels,
  743. kernel_size=3,
  744. stride=1,
  745. padding=1)
  746. def forward(self, x, orig_zoom):
  747. #orig_zoom is a list of the zooms [zh, zw, zz]
  748. #assert x.shape[2] == x.shape[3] == self.resolution, "{}, {}, {}".format(x.shape[2], x.shape[3], self.resolution)
  749. # timestep embedding
  750. temb = None
  751. shapes = x.shape
  752. h = shapes[2]
  753. w = shapes[3]
  754. output_shape = (h, w)
  755. # downsampling
  756. hin = self.conv_in(x) #hin are tensors of shape [1, 128, 256, 256]
  757. hin, inner_network_shape_perbs, padding_info = rescale_to_uniform_voxel_size(hin, orig_zoom, iz_h = self.inner_res, iz_w = self.inner_res)
  758. hs = [self.conv_vinn(hin)]
  759. for i_level in range(self.num_resolutions): #num of resolutions is 5
  760. for i_block in range(self.num_res_blocks):# i_block iterates from 0 to num of resolutions (5)
  761. h = self.down[i_level].block[i_block](hs[-1], temb)
  762. if len(self.down[i_level].attn) > 0:
  763. h = self.down[i_level].attn[i_block](h)
  764. hs.append(h)
  765. if i_level != self.num_resolutions-1:
  766. hs.append(self.down[i_level].downsample(hs[-1]))
  767. # middle
  768. h = hs[-1]
  769. h = self.mid.block_1(h, temb)
  770. h = self.mid.attn_1(h)
  771. h = self.mid.block_2(h, temb)
  772. # end
  773. h = self.norm_out(h)
  774. h = nonlinearity(h)
  775. h = self.conv_out(h)
  776. return h, inner_network_shape_perbs, padding_info, output_shape
  777. class DecoderVINN(nn.Module):
  778. def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
  779. attn_resolutions, dropout=0.0, resamp_with_conv=True,
  780. inner_res = 0.7,
  781. in_channels, z_channels, give_pre_end=False, **ignorekwargs):
  782. super().__init__()
  783. self.ch = ch
  784. self.temb_ch = 0
  785. self.num_resolutions = len(ch_mult)
  786. self.num_res_blocks = num_res_blocks
  787. self.in_channels = in_channels
  788. self.give_pre_end = give_pre_end
  789. self.inner_res = inner_res
  790. # compute in_ch_mult, block_in and curr_res at lowest res
  791. in_ch_mult = (1,)+tuple(ch_mult)
  792. block_in = ch*ch_mult[self.num_resolutions-1]
  793. curr_res = inner_res * 2**(self.num_resolutions-1)
  794. self.z_shape = (1,z_channels,curr_res,curr_res)
  795. print("Working with z of shape = {} mm.".format(self.z_shape))
  796. print(f'The internal resolution of the network (VINN) is: {inner_res} mm isotropic ..')
  797. # z to block_in
  798. self.conv_in = torch.nn.Conv2d(z_channels,
  799. block_in,
  800. kernel_size=3,
  801. stride=1,
  802. padding=1)
  803. # middle
  804. self.mid = nn.Module()
  805. self.mid.block_1 = ResnetBlock(in_channels=block_in,
  806. out_channels=block_in,
  807. temb_channels=self.temb_ch,
  808. dropout=dropout)
  809. self.mid.attn_1 = AttnBlock(block_in)
  810. self.mid.block_2 = ResnetBlock(in_channels=block_in,
  811. out_channels=block_in,
  812. temb_channels=self.temb_ch,
  813. dropout=dropout)
  814. # upsampling
  815. self.up = nn.ModuleList()
  816. for i_level in reversed(range(self.num_resolutions)):
  817. block = nn.ModuleList()
  818. attn = nn.ModuleList()
  819. block_out = ch*ch_mult[i_level]
  820. for i_block in range(self.num_res_blocks+1):
  821. block.append(ResnetBlock(in_channels=block_in,
  822. out_channels=block_out,
  823. temb_channels=self.temb_ch,
  824. dropout=dropout))
  825. block_in = block_out
  826. if curr_res in attn_resolutions:
  827. attn.append(AttnBlock(block_in))
  828. up = nn.Module()
  829. up.block = block
  830. up.attn = attn
  831. if i_level != 0:
  832. up.upsample = Upsample(block_in, resamp_with_conv)
  833. curr_res = curr_res * 2
  834. self.up.insert(0, up) # prepend to get consistent order
  835. # end
  836. self.norm_out = Normalize(block_in)
  837. self.conv_out_vinn = torch.nn.Conv2d(block_in,
  838. block_in,
  839. kernel_size=3,
  840. stride=1,
  841. padding=1)
  842. self.conv_out = torch.nn.Conv2d(block_in,
  843. out_ch,
  844. kernel_size=3,
  845. stride=1,
  846. padding=1)
  847. def forward(self, z, inner_network_shape_perbs, padding_info, output_shape):
  848. #assert z.shape[1:] == self.z_shape[1:]
  849. self.last_z_shape = z.shape
  850. # timestep embedding
  851. temb = None
  852. # z to block_in
  853. h = self.conv_in(z)
  854. # middle
  855. h = self.mid.block_1(h, temb)
  856. h = self.mid.attn_1(h)
  857. h = self.mid.block_2(h, temb)
  858. # upsampling
  859. for i_level in reversed(range(self.num_resolutions)):
  860. for i_block in range(self.num_res_blocks+1):
  861. h = self.up[i_level].block[i_block](h, temb)
  862. if len(self.up[i_level].attn) > 0:
  863. h = self.up[i_level].attn[i_block](h)
  864. if i_level != 0:
  865. h = self.up[i_level].upsample(h)
  866. # end
  867. if self.give_pre_end:
  868. return h
  869. h = self.norm_out(h)
  870. h = nonlinearity(h)
  871. #Bringing feature maps to native space:
  872. h = invert_rescale_to_uniform_voxel_size(h, inner_network_shape_perbs, padding_info, output_shape)
  873. h = self.conv_out_vinn(h)
  874. h = self.conv_out(h)
  875. return h
  876. class Encoder(nn.Module):
  877. def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
  878. attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels,
  879. resolution, z_channels, double_z=True, **ignore_kwargs):
  880. super().__init__()
  881. self.ch = ch
  882. self.temb_ch = 0
  883. self.num_resolutions = len(ch_mult)
  884. self.num_res_blocks = num_res_blocks
  885. self.resolution = resolution
  886. self.in_channels = in_channels
  887. # downsampling
  888. self.conv_in = torch.nn.Conv2d(in_channels,
  889. self.ch,
  890. kernel_size=3,
  891. stride=1,
  892. padding=1)
  893. curr_res = resolution
  894. in_ch_mult = (1,)+tuple(ch_mult)
  895. self.down = nn.ModuleList()
  896. for i_level in range(self.num_resolutions):
  897. block = nn.ModuleList()
  898. attn = nn.ModuleList()
  899. block_in = ch*in_ch_mult[i_level]
  900. block_out = ch*ch_mult[i_level]
  901. for i_block in range(self.num_res_blocks):
  902. block.append(ResnetBlock(in_channels=block_in,
  903. out_channels=block_out,
  904. temb_channels=self.temb_ch,
  905. dropout=dropout))
  906. block_in = block_out
  907. if curr_res in attn_resolutions:
  908. attn.append(AttnBlock(block_in))
  909. down = nn.Module()
  910. down.block = block
  911. down.attn = attn
  912. if i_level != self.num_resolutions-1:
  913. down.downsample = Downsample(block_in, resamp_with_conv)
  914. curr_res = curr_res // 2
  915. self.down.append(down)
  916. # middle
  917. self.mid = nn.Module()
  918. self.mid.block_1 = ResnetBlock(in_channels=block_in,
  919. out_channels=block_in,
  920. temb_channels=self.temb_ch,
  921. dropout=dropout)
  922. self.mid.attn_1 = AttnBlock(block_in)
  923. self.mid.block_2 = ResnetBlock(in_channels=block_in,
  924. out_channels=block_in,
  925. temb_channels=self.temb_ch,
  926. dropout=dropout)
  927. # end
  928. self.norm_out = Normalize(block_in)
  929. self.conv_out = torch.nn.Conv2d(block_in,
  930. 2*z_channels if double_z else z_channels,
  931. kernel_size=3,
  932. stride=1,
  933. padding=1)
  934. def forward(self, x):
  935. #assert x.shape[2] == x.shape[3] == self.resolution, "{}, {}, {}".format(x.shape[2], x.shape[3], self.resolution)
  936. # timestep embedding
  937. temb = None
  938. # downsampling
  939. hs = [self.conv_in(x)] #hs is a list of tensors of shape [1, 128, 256, 256]
  940. for i_level in range(self.num_resolutions): #num of resolutions is 5
  941. for i_block in range(self.num_res_blocks):# i_block iterates from 0 to num of resolutions (5)
  942. h = self.down[i_level].block[i_block](hs[-1], temb)
  943. if len(self.down[i_level].attn) > 0:
  944. h = self.down[i_level].attn[i_block](h)
  945. hs.append(h)
  946. if i_level != self.num_resolutions-1:
  947. hs.append(self.down[i_level].downsample(hs[-1]))
  948. # middle
  949. h = hs[-1]
  950. h = self.mid.block_1(h, temb)
  951. h = self.mid.attn_1(h)
  952. h = self.mid.block_2(h, temb)
  953. # end
  954. h = self.norm_out(h)
  955. h = nonlinearity(h)
  956. h = self.conv_out(h)
  957. return h
  958. class Decoder(nn.Module):
  959. def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
  960. attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels,
  961. resolution, z_channels, give_pre_end=False, **ignorekwargs):
  962. super().__init__()
  963. self.ch = ch
  964. self.temb_ch = 0
  965. self.num_resolutions = len(ch_mult)
  966. self.num_res_blocks = num_res_blocks
  967. self.resolution = resolution
  968. self.in_channels = in_channels
  969. self.give_pre_end = give_pre_end
  970. # compute in_ch_mult, block_in and curr_res at lowest res
  971. in_ch_mult = (1,)+tuple(ch_mult)
  972. block_in = ch*ch_mult[self.num_resolutions-1]
  973. curr_res = resolution // 2**(self.num_resolutions-1)
  974. self.z_shape = (1,z_channels,curr_res,curr_res)
  975. print("Working with z of shape {} = {} dimensions.".format(
  976. self.z_shape, np.prod(self.z_shape)))
  977. # z to block_in
  978. self.conv_in = torch.nn.Conv2d(z_channels,
  979. block_in,
  980. kernel_size=3,
  981. stride=1,
  982. padding=1)
  983. # middle
  984. self.mid = nn.Module()
  985. self.mid.block_1 = ResnetBlock(in_channels=block_in,
  986. out_channels=block_in,
  987. temb_channels=self.temb_ch,
  988. dropout=dropout)
  989. self.mid.attn_1 = AttnBlock(block_in)
  990. self.mid.block_2 = ResnetBlock(in_channels=block_in,
  991. out_channels=block_in,
  992. temb_channels=self.temb_ch,
  993. dropout=dropout)
  994. # upsampling
  995. self.up = nn.ModuleList()
  996. for i_level in reversed(range(self.num_resolutions)):
  997. block = nn.ModuleList()
  998. attn = nn.ModuleList()
  999. block_out = ch*ch_mult[i_level]
  1000. for i_block in range(self.num_res_blocks+1):
  1001. block.append(ResnetBlock(in_channels=block_in,
  1002. out_channels=block_out,
  1003. temb_channels=self.temb_ch,
  1004. dropout=dropout))
  1005. block_in = block_out
  1006. if curr_res in attn_resolutions:
  1007. attn.append(AttnBlock(block_in))
  1008. up = nn.Module()
  1009. up.block = block
  1010. up.attn = attn
  1011. if i_level != 0:
  1012. up.upsample = Upsample(block_in, resamp_with_conv)
  1013. curr_res = curr_res * 2
  1014. self.up.insert(0, up) # prepend to get consistent order
  1015. # end
  1016. self.norm_out = Normalize(block_in)
  1017. self.conv_out = torch.nn.Conv2d(block_in,
  1018. out_ch,
  1019. kernel_size=3,
  1020. stride=1,
  1021. padding=1)
  1022. def forward(self, z):
  1023. #assert z.shape[1:] == self.z_shape[1:]
  1024. self.last_z_shape = z.shape
  1025. # timestep embedding
  1026. temb = None
  1027. # z to block_in
  1028. h = self.conv_in(z)
  1029. # middle
  1030. h = self.mid.block_1(h, temb)
  1031. h = self.mid.attn_1(h)
  1032. h = self.mid.block_2(h, temb)
  1033. # upsampling
  1034. for i_level in reversed(range(self.num_resolutions)):
  1035. for i_block in range(self.num_res_blocks+1):
  1036. h = self.up[i_level].block[i_block](h, temb)
  1037. if len(self.up[i_level].attn) > 0:
  1038. h = self.up[i_level].attn[i_block](h)
  1039. if i_level != 0:
  1040. h = self.up[i_level].upsample(h)
  1041. # end
  1042. if self.give_pre_end:
  1043. return h
  1044. h = self.norm_out(h)
  1045. h = nonlinearity(h)
  1046. h = self.conv_out(h)
  1047. return h
  1048. class VUNet(nn.Module):
  1049. def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
  1050. attn_resolutions, dropout=0.0, resamp_with_conv=True,
  1051. in_channels, c_channels,
  1052. resolution, z_channels, use_timestep=False, **ignore_kwargs):
  1053. super().__init__()
  1054. self.ch = ch
  1055. self.temb_ch = self.ch*4
  1056. self.num_resolutions = len(ch_mult)
  1057. self.num_res_blocks = num_res_blocks
  1058. self.resolution = resolution
  1059. self.use_timestep = use_timestep
  1060. if self.use_timestep:
  1061. # timestep embedding
  1062. self.temb = nn.Module()
  1063. self.temb.dense = nn.ModuleList([
  1064. torch.nn.Linear(self.ch,
  1065. self.temb_ch),
  1066. torch.nn.Linear(self.temb_ch,
  1067. self.temb_ch),
  1068. ])
  1069. # downsampling
  1070. self.conv_in = torch.nn.Conv2d(c_channels,
  1071. self.ch,
  1072. kernel_size=3,
  1073. stride=1,
  1074. padding=1)
  1075. curr_res = resolution
  1076. in_ch_mult = (1,)+tuple(ch_mult)
  1077. self.down = nn.ModuleList()
  1078. for i_level in range(self.num_resolutions):
  1079. block = nn.ModuleList()
  1080. attn = nn.ModuleList()
  1081. block_in = ch*in_ch_mult[i_level]
  1082. block_out = ch*ch_mult[i_level]
  1083. for i_block in range(self.num_res_blocks):
  1084. block.append(ResnetBlock(in_channels=block_in,
  1085. out_channels=block_out,
  1086. temb_channels=self.temb_ch,
  1087. dropout=dropout))
  1088. block_in = block_out
  1089. if curr_res in attn_resolutions:
  1090. attn.append(AttnBlock(block_in))
  1091. down = nn.Module()
  1092. down.block = block
  1093. down.attn = attn
  1094. if i_level != self.num_resolutions-1:
  1095. down.downsample = Downsample(block_in, resamp_with_conv)
  1096. curr_res = curr_res // 2
  1097. self.down.append(down)
  1098. self.z_in = torch.nn.Conv2d(z_channels,
  1099. block_in,
  1100. kernel_size=1,
  1101. stride=1,
  1102. padding=0)
  1103. # middle
  1104. self.mid = nn.Module()
  1105. self.mid.block_1 = ResnetBlock(in_channels=2*block_in,
  1106. out_channels=block_in,
  1107. temb_channels=self.temb_ch,
  1108. dropout=dropout)
  1109. self.mid.attn_1 = AttnBlock(block_in)
  1110. self.mid.block_2 = ResnetBlock(in_channels=block_in,
  1111. out_channels=block_in,
  1112. temb_channels=self.temb_ch,
  1113. dropout=dropout)
  1114. # upsampling
  1115. self.up = nn.ModuleList()
  1116. for i_level in reversed(range(self.num_resolutions)):
  1117. block = nn.ModuleList()
  1118. attn = nn.ModuleList()
  1119. block_out = ch*ch_mult[i_level]
  1120. skip_in = ch*ch_mult[i_level]
  1121. for i_block in range(self.num_res_blocks+1):
  1122. if i_block == self.num_res_blocks:
  1123. skip_in = ch*in_ch_mult[i_level]
  1124. block.append(ResnetBlock(in_channels=block_in+skip_in,
  1125. out_channels=block_out,
  1126. temb_channels=self.temb_ch,
  1127. dropout=dropout))
  1128. block_in = block_out
  1129. if curr_res in attn_resolutions:
  1130. attn.append(AttnBlock(block_in))
  1131. up = nn.Module()
  1132. up.block = block
  1133. up.attn = attn
  1134. if i_level != 0:
  1135. up.upsample = Upsample(block_in, resamp_with_conv)
  1136. curr_res = curr_res * 2
  1137. self.up.insert(0, up) # prepend to get consistent order
  1138. # end
  1139. self.norm_out = Normalize(block_in)
  1140. self.conv_out = torch.nn.Conv2d(block_in,
  1141. out_ch,
  1142. kernel_size=3,
  1143. stride=1,
  1144. padding=1)
  1145. def forward(self, x, z):
  1146. #assert x.shape[2] == x.shape[3] == self.resolution
  1147. if self.use_timestep:
  1148. # timestep embedding
  1149. assert t is not None
  1150. temb = get_timestep_embedding(t, self.ch)
  1151. temb = self.temb.dense[0](temb)
  1152. temb = nonlinearity(temb)
  1153. temb = self.temb.dense[1](temb)
  1154. else:
  1155. temb = None
  1156. # downsampling
  1157. hs = [self.conv_in(x)]
  1158. for i_level in range(self.num_resolutions):
  1159. for i_block in range(self.num_res_blocks):
  1160. h = self.down[i_level].block[i_block](hs[-1], temb)
  1161. if len(self.down[i_level].attn) > 0:
  1162. h = self.down[i_level].attn[i_block](h)
  1163. hs.append(h)
  1164. if i_level != self.num_resolutions-1:
  1165. hs.append(self.down[i_level].downsample(hs[-1]))
  1166. # middle
  1167. h = hs[-1]
  1168. z = self.z_in(z)
  1169. h = torch.cat((h,z),dim=1)
  1170. h = self.mid.block_1(h, temb)
  1171. h = self.mid.attn_1(h)
  1172. h = self.mid.block_2(h, temb)
  1173. # upsampling
  1174. for i_level in reversed(range(self.num_resolutions)):
  1175. for i_block in range(self.num_res_blocks+1):
  1176. h = self.up[i_level].block[i_block](
  1177. torch.cat([h, hs.pop()], dim=1), temb)
  1178. if len(self.up[i_level].attn) > 0:
  1179. h = self.up[i_level].attn[i_block](h)
  1180. if i_level != 0:
  1181. h = self.up[i_level].upsample(h)
  1182. # end
  1183. h = self.norm_out(h)
  1184. h = nonlinearity(h)
  1185. h = self.conv_out(h)
  1186. return h
  1187. class SimpleDecoder(nn.Module):
  1188. def __init__(self, in_channels, out_channels, *args, **kwargs):
  1189. super().__init__()
  1190. self.model = nn.ModuleList([nn.Conv2d(in_channels, in_channels, 1),
  1191. ResnetBlock(in_channels=in_channels,
  1192. out_channels=2 * in_channels,
  1193. temb_channels=0, dropout=0.0),
  1194. ResnetBlock(in_channels=2 * in_channels,
  1195. out_channels=4 * in_channels,
  1196. temb_channels=0, dropout=0.0),
  1197. ResnetBlock(in_channels=4 * in_channels,
  1198. out_channels=2 * in_channels,
  1199. temb_channels=0, dropout=0.0),
  1200. nn.Conv2d(2*in_channels, in_channels, 1),
  1201. Upsample(in_channels, with_conv=True)])
  1202. # end
  1203. self.norm_out = Normalize(in_channels)
  1204. self.conv_out = torch.nn.Conv2d(in_channels,
  1205. out_channels,
  1206. kernel_size=3,
  1207. stride=1,
  1208. padding=1)
  1209. def forward(self, x):
  1210. for i, layer in enumerate(self.model):
  1211. if i in [1,2,3]:
  1212. x = layer(x, None)
  1213. else:
  1214. x = layer(x)
  1215. h = self.norm_out(x)
  1216. h = nonlinearity(h)
  1217. x = self.conv_out(h)
  1218. return x
  1219. class UpsampleDecoder(nn.Module):
  1220. def __init__(self, in_channels, out_channels, ch, num_res_blocks, resolution,
  1221. ch_mult=(2,2), dropout=0.0):
  1222. super().__init__()
  1223. # upsampling
  1224. self.temb_ch = 0
  1225. self.num_resolutions = len(ch_mult)
  1226. self.num_res_blocks = num_res_blocks
  1227. block_in = in_channels
  1228. curr_res = resolution // 2 ** (self.num_resolutions - 1)
  1229. self.res_blocks = nn.ModuleList()
  1230. self.upsample_blocks = nn.ModuleList()
  1231. for i_level in range(self.num_resolutions):
  1232. res_block = []
  1233. block_out = ch * ch_mult[i_level]
  1234. for i_block in range(self.num_res_blocks + 1):
  1235. res_block.append(ResnetBlock(in_channels=block_in,
  1236. out_channels=block_out,
  1237. temb_channels=self.temb_ch,
  1238. dropout=dropout))
  1239. block_in = block_out
  1240. self.res_blocks.append(nn.ModuleList(res_block))
  1241. if i_level != self.num_resolutions - 1:
  1242. self.upsample_blocks.append(Upsample(block_in, True))
  1243. curr_res = curr_res * 2
  1244. # end
  1245. self.norm_out = Normalize(block_in)
  1246. self.conv_out = torch.nn.Conv2d(block_in,
  1247. out_channels,
  1248. kernel_size=3,
  1249. stride=1,
  1250. padding=1)
  1251. def forward(self, x):
  1252. # upsampling
  1253. h = x
  1254. for k, i_level in enumerate(range(self.num_resolutions)):
  1255. for i_block in range(self.num_res_blocks + 1):
  1256. h = self.res_blocks[i_level][i_block](h, None)
  1257. if i_level != self.num_resolutions - 1:
  1258. h = self.upsample_blocks[k](h)
  1259. h = self.norm_out(h)
  1260. h = nonlinearity(h)
  1261. h = self.conv_out(h)
  1262. return h
  1263. class Encoder3D(nn.Module):
  1264. def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
  1265. attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels,
  1266. resolution, z_channels, double_z=True, **ignore_kwargs):
  1267. super().__init__()
  1268. self.ch = ch
  1269. self.temb_ch = 0
  1270. self.num_resolutions = len(ch_mult)
  1271. self.num_res_blocks = num_res_blocks
  1272. self.resolution = resolution
  1273. self.in_channels = in_channels
  1274. # downsampling
  1275. self.conv_in = torch.nn.Conv3d(in_channels,
  1276. self.ch,
  1277. kernel_size=3,
  1278. stride=1,
  1279. padding=1)
  1280. curr_res = resolution
  1281. in_ch_mult = (1,)+tuple(ch_mult)
  1282. self.down = nn.ModuleList()
  1283. for i_level in range(self.num_resolutions):
  1284. block = nn.ModuleList()
  1285. attn = nn.ModuleList()
  1286. block_in = ch*in_ch_mult[i_level]
  1287. block_out = ch*ch_mult[i_level]
  1288. for i_block in range(self.num_res_blocks):
  1289. block.append(ResnetBlock3D(in_channels=block_in,
  1290. out_channels=block_out,
  1291. temb_channels=self.temb_ch,
  1292. dropout=dropout))
  1293. block_in = block_out
  1294. if curr_res in attn_resolutions:
  1295. attn.append(AttnBlock3D(block_in))
  1296. down = nn.Module()
  1297. down.block = block
  1298. down.attn = attn
  1299. if i_level != self.num_resolutions-1:
  1300. down.downsample = Downsample3D(block_in, resamp_with_conv)
  1301. curr_res = curr_res // 2
  1302. self.down.append(down)
  1303. # middle
  1304. self.mid = nn.Module()
  1305. self.mid.block_1 = ResnetBlock3D(in_channels=block_in,
  1306. out_channels=block_in,
  1307. temb_channels=self.temb_ch,
  1308. dropout=dropout)
  1309. self.mid.attn_1 = AttnBlock3D(block_in)
  1310. self.mid.block_2 = ResnetBlock3D(in_channels=block_in,
  1311. out_channels=block_in,
  1312. temb_channels=self.temb_ch,
  1313. dropout=dropout)
  1314. # end
  1315. self.norm_out = Normalize(block_in)
  1316. self.conv_out = torch.nn.Conv3d(block_in,
  1317. 2*z_channels if double_z else z_channels,
  1318. kernel_size=3,
  1319. stride=1,
  1320. padding=1)
  1321. def forward(self, x):
  1322. #assert x.shape[2] == x.shape[3] == self.resolution, "{}, {}, {}".format(x.shape[2], x.shape[3], self.resolution)
  1323. # timestep embedding
  1324. temb = None
  1325. # downsampling
  1326. hs = [self.conv_in(x)] #3D Input of shape expected [BS, C, D, H, W]
  1327. for i_level in range(self.num_resolutions): #num of resolutions is 5
  1328. for i_block in range(self.num_res_blocks):# i_block iterates from 0 to num of resolutions (5)
  1329. h = self.down[i_level].block[i_block](hs[-1], temb)
  1330. if len(self.down[i_level].attn) > 0:
  1331. h = self.down[i_level].attn[i_block](h)
  1332. hs.append(h)
  1333. if i_level != self.num_resolutions-1:
  1334. hs.append(self.down[i_level].downsample(hs[-1]))
  1335. # middle
  1336. h = hs[-1]
  1337. h = self.mid.block_1(h, temb)
  1338. h = self.mid.attn_1(h)
  1339. h = self.mid.block_2(h, temb)
  1340. # end
  1341. h = self.norm_out(h)
  1342. h = nonlinearity(h)
  1343. h = self.conv_out(h) #Shape torch.Size([1, 2, 48, 48, 8]) GPU: 11133MiB / 23028MiB
  1344. return h
  1345. class Decoder3D(nn.Module):
  1346. def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
  1347. attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels,
  1348. resolution, z_channels, give_pre_end=False, **ignorekwargs):
  1349. super().__init__()
  1350. self.ch = ch
  1351. self.temb_ch = 0
  1352. self.num_resolutions = len(ch_mult)
  1353. self.num_res_blocks = num_res_blocks
  1354. self.resolution = resolution
  1355. self.in_channels = in_channels
  1356. self.give_pre_end = give_pre_end
  1357. # compute in_ch_mult, block_in and curr_res at lowest res
  1358. in_ch_mult = (1,)+tuple(ch_mult)
  1359. block_in = ch*ch_mult[self.num_resolutions-1]
  1360. curr_res = resolution // 2**(self.num_resolutions-1)
  1361. self.z_shape = (1,z_channels,curr_res,curr_res)
  1362. print("Working with z of shape {} = {} dimensions.".format(
  1363. self.z_shape, np.prod(self.z_shape)))
  1364. # z to block_in
  1365. self.conv_in = torch.nn.Conv3d(z_channels,
  1366. block_in,
  1367. kernel_size=3,
  1368. stride=1,
  1369. padding=1)
  1370. # middle
  1371. self.mid = nn.Module()
  1372. self.mid.block_1 = ResnetBlock3D(in_channels=block_in,
  1373. out_channels=block_in,
  1374. temb_channels=self.temb_ch,
  1375. dropout=dropout)
  1376. self.mid.attn_1 = AttnBlock3D(block_in)
  1377. self.mid.block_2 = ResnetBlock3D(in_channels=block_in,
  1378. out_channels=block_in,
  1379. temb_channels=self.temb_ch,
  1380. dropout=dropout)
  1381. # upsampling
  1382. self.up = nn.ModuleList()
  1383. for i_level in reversed(range(self.num_resolutions)):
  1384. block = nn.ModuleList()
  1385. attn = nn.ModuleList()
  1386. block_out = ch*ch_mult[i_level]
  1387. for i_block in range(self.num_res_blocks+1):
  1388. block.append(ResnetBlock3D(in_channels=block_in,
  1389. out_channels=block_out,
  1390. temb_channels=self.temb_ch,
  1391. dropout=dropout))
  1392. block_in = block_out
  1393. if curr_res in attn_resolutions:
  1394. attn.append(AttnBlock3D(block_in))
  1395. up = nn.Module()
  1396. up.block = block
  1397. up.attn = attn
  1398. if i_level != 0:
  1399. up.upsample = Upsample3D(block_in, resamp_with_conv)
  1400. curr_res = curr_res * 2
  1401. self.up.insert(0, up) # prepend to get consistent order
  1402. # end
  1403. self.norm_out = Normalize(block_in)
  1404. self.conv_out = torch.nn.Conv3d(block_in,
  1405. out_ch,
  1406. kernel_size=3,
  1407. stride=1,
  1408. padding=1)
  1409. def forward(self, z):
  1410. #assert z.shape[1:] == self.z_shape[1:]
  1411. self.last_z_shape = z.shape
  1412. # timestep embedding
  1413. temb = None
  1414. # z to block_in
  1415. h = self.conv_in(z)
  1416. # middle
  1417. h = self.mid.block_1(h, temb)
  1418. h = self.mid.attn_1(h)
  1419. h = self.mid.block_2(h, temb)
  1420. # upsampling
  1421. for i_level in reversed(range(self.num_resolutions)):
  1422. for i_block in range(self.num_res_blocks+1):
  1423. h = self.up[i_level].block[i_block](h, temb)
  1424. if len(self.up[i_level].attn) > 0:
  1425. h = self.up[i_level].attn[i_block](h)
  1426. if i_level != 0:
  1427. h = self.up[i_level].upsample(h)
  1428. # end
  1429. if self.give_pre_end:
  1430. return h
  1431. h = self.norm_out(h)
  1432. h = nonlinearity(h)
  1433. h = self.conv_out(h)
  1434. return h
  1435. import torch
  1436. import torch.nn as nn
  1437. import torch.nn.functional as F
  1438. import numpy as np
  1439. # Assuming Normalize and nonlinearity are defined earlier in the file
  1440. # from .model import Normalize, nonlinearity # Or similar import
  1441. class ResnetBlock3DFiLM(nn.Module):
  1442. """
  1443. 3D ResNet block modified to incorporate FiLM conditioning. (Corrected)
  1444. Uses nn.SiLU() in Sequential, nonlinearity() in forward.
  1445. """
  1446. def __init__(self, *, in_channels, out_channels=None, conv_shortcut=False,
  1447. dropout, temb_channels=512, cond_embed_dim=None):
  1448. super().__init__()
  1449. self.in_channels = in_channels
  1450. out_channels = in_channels if out_channels is None else out_channels
  1451. self.out_channels = out_channels
  1452. self.use_conv_shortcut = conv_shortcut
  1453. self.cond_embed_dim = cond_embed_dim
  1454. self.norm1 = Normalize(in_channels)
  1455. self.conv1 = nn.Conv3d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
  1456. if temb_channels is not None and temb_channels > 0:
  1457. self.temb_proj = nn.Linear(temb_channels, out_channels)
  1458. else:
  1459. self.temb_proj = None # Explicitly set to None if not used
  1460. self.norm2 = Normalize(out_channels)
  1461. self.dropout = nn.Dropout(dropout)
  1462. self.conv2 = nn.Conv3d(out_channels, out_channels, kernel_size=3, stride=1, padding=1)
  1463. # FiLM generation layer
  1464. if self.cond_embed_dim is not None and self.cond_embed_dim > 0:
  1465. self.cond_proj = nn.Sequential(
  1466. nn.SiLU(), # Correct module for Sequential
  1467. nn.Linear(self.cond_embed_dim, 2 * self.out_channels)
  1468. )
  1469. if len(self.cond_proj) > 1 and isinstance(self.cond_proj[-1], nn.Linear):
  1470. nn.init.zeros_(self.cond_proj[-1].weight)
  1471. nn.init.zeros_(self.cond_proj[-1].bias)
  1472. else:
  1473. self.cond_proj = None # Ensure it's None if not used
  1474. # Shortcut connection
  1475. if self.in_channels != self.out_channels:
  1476. if self.use_conv_shortcut:
  1477. self.conv_shortcut = nn.Conv3d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
  1478. else:
  1479. self.nin_shortcut = nn.Conv3d(in_channels, out_channels, kernel_size=1, stride=1, padding=0)
  1480. else:
  1481. self.conv_shortcut = None
  1482. self.nin_shortcut = None
  1483. def forward(self, x, temb, cond_emb=None):
  1484. h = x
  1485. h = self.conv1( nonlinearity( self.norm1(h) ) )
  1486. if temb is not None and self.temb_proj is not None:
  1487. h = h + self.temb_proj(nonlinearity(temb))[:, :, None, None, None]
  1488. h_norm = self.norm2(h)
  1489. if cond_emb is not None and self.cond_proj is not None:
  1490. gamma_beta = self.cond_proj(cond_emb)[:, :, None, None, None]
  1491. gamma, beta = torch.chunk(gamma_beta, 2, dim=1)
  1492. h_modulated = h_norm * (1. + gamma) + beta
  1493. else:
  1494. h_modulated = h_norm
  1495. h = self.conv2( self.dropout( nonlinearity(h_modulated) ) )
  1496. if self.in_channels != self.out_channels:
  1497. if self.use_conv_shortcut:
  1498. x_shortcut = self.conv_shortcut(x)
  1499. elif self.nin_shortcut is not None:
  1500. x_shortcut = self.nin_shortcut(x)
  1501. else: # Should not happen if in_channels != out_channels, but safety
  1502. x_shortcut = x
  1503. else:
  1504. x_shortcut = x
  1505. return x_shortcut + h
  1506. class Encoder3D_aniso(nn.Module):
  1507. def __init__(self, *, ch, out_ch, ch_mult=(1, 2, 4), num_res_blocks=2,
  1508. attn_resolutions=(), dropout=0.0, resamp_with_conv=True, in_channels=1,
  1509. resolution=192, depth=32, z_channels=1, double_z=True, **ignore_kwargs):
  1510. super().__init__()
  1511. self.ch = ch
  1512. self.temb_ch = 0
  1513. self.num_resolutions = len(ch_mult)
  1514. self.num_res_blocks = num_res_blocks
  1515. self.resolution = resolution
  1516. self.depth = depth
  1517. self.in_channels = in_channels
  1518. # initial conv
  1519. self.conv_in = nn.Conv3d(in_channels, ch, kernel_size=3, stride=1, padding=1)
  1520. curr_res_h = resolution
  1521. curr_res_w = resolution
  1522. curr_res_d = depth
  1523. in_ch_mult = (1,) + tuple(ch_mult)
  1524. self.down = nn.ModuleList()
  1525. for i_level in range(self.num_resolutions):
  1526. block = nn.ModuleList()
  1527. attn = nn.ModuleList()
  1528. block_in = ch * in_ch_mult[i_level]
  1529. block_out = ch * ch_mult[i_level]
  1530. for _ in range(self.num_res_blocks):
  1531. block.append(ResnetBlock3D(in_channels=block_in,
  1532. out_channels=block_out,
  1533. temb_channels=self.temb_ch,
  1534. dropout=dropout))
  1535. block_in = block_out
  1536. if curr_res_h in attn_resolutions:
  1537. attn.append(AttnBlock3D(block_in))
  1538. down = nn.Module()
  1539. down.block = block
  1540. down.attn = attn
  1541. if i_level != self.num_resolutions - 1:
  1542. down.downsample = Downsample3D_HW(block_in, resamp_with_conv)
  1543. curr_res_h //= 2
  1544. curr_res_w //= 2
  1545. self.down.append(down)
  1546. # middle
  1547. self.mid = nn.Module()
  1548. self.mid.block_1 = ResnetBlock3D(in_channels=block_in, out_channels=block_in,
  1549. temb_channels=self.temb_ch, dropout=dropout)
  1550. self.mid.attn_1 = AttnBlock3D(block_in)
  1551. self.mid.block_2 = ResnetBlock3D(in_channels=block_in, out_channels=block_in,
  1552. temb_channels=self.temb_ch, dropout=dropout)
  1553. # end
  1554. self.norm_out = Normalize(block_in)
  1555. self.conv_out = nn.Conv3d(block_in, 2 * z_channels if double_z else z_channels,
  1556. kernel_size=3, stride=1, padding=1)
  1557. def forward(self, x):
  1558. temb = None
  1559. hs = [self.conv_in(x)]
  1560. for i_level in range(self.num_resolutions):
  1561. for i_block in range(self.num_res_blocks):
  1562. h = self.down[i_level].block[i_block](hs[-1], temb)
  1563. if len(self.down[i_level].attn) > 0:
  1564. h = self.down[i_level].attn[i_block](h)
  1565. hs.append(h)
  1566. if i_level != self.num_resolutions - 1:
  1567. hs.append(self.down[i_level].downsample(hs[-1]))
  1568. h = hs[-1]
  1569. h = self.mid.block_1(h, temb)
  1570. h = self.mid.attn_1(h)
  1571. h = self.mid.block_2(h, temb)
  1572. h = self.norm_out(h)
  1573. h = nonlinearity(h)
  1574. h = self.conv_out(h)
  1575. return h
  1576. class Decoder3D_aniso(nn.Module):
  1577. def __init__(self, *, ch, out_ch, ch_mult=(1, 2, 4), num_res_blocks=2,
  1578. attn_resolutions=(), dropout=0.0, resamp_with_conv=True, in_channels=1,
  1579. resolution=192, depth=32, z_channels=1, give_pre_end=False, **ignore_kwargs):
  1580. super().__init__()
  1581. self.ch = ch
  1582. self.temb_ch = 0
  1583. self.num_resolutions = len(ch_mult)
  1584. self.num_res_blocks = num_res_blocks
  1585. self.resolution = resolution
  1586. self.depth = depth
  1587. self.in_channels = in_channels
  1588. self.give_pre_end = give_pre_end
  1589. in_ch_mult = (1,) + tuple(ch_mult)
  1590. block_in = ch * ch_mult[self.num_resolutions - 1]
  1591. curr_h = resolution // 2 ** (self.num_resolutions - 1)
  1592. curr_w = resolution // 2 ** (self.num_resolutions - 1)
  1593. curr_d = depth # unchanged
  1594. self.z_shape = (1, z_channels, curr_d, curr_h, curr_w)
  1595. # z to block_in
  1596. self.conv_in = nn.Conv3d(z_channels, block_in, kernel_size=3, stride=1, padding=1)
  1597. # middle
  1598. self.mid = nn.Module()
  1599. self.mid.block_1 = ResnetBlock3D(in_channels=block_in, out_channels=block_in,
  1600. temb_channels=self.temb_ch, dropout=dropout)
  1601. self.mid.attn_1 = AttnBlock3D(block_in)
  1602. self.mid.block_2 = ResnetBlock3D(in_channels=block_in, out_channels=block_in,
  1603. temb_channels=self.temb_ch, dropout=dropout)
  1604. # upsampling
  1605. self.up = nn.ModuleList()
  1606. for i_level in reversed(range(self.num_resolutions)):
  1607. block = nn.ModuleList()
  1608. attn = nn.ModuleList()
  1609. block_out = ch * ch_mult[i_level]
  1610. for _ in range(self.num_res_blocks + 1):
  1611. block.append(ResnetBlock3D(in_channels=block_in, out_channels=block_out,
  1612. temb_channels=self.temb_ch, dropout=dropout))
  1613. block_in = block_out
  1614. if (resolution // 2 ** i_level) in attn_resolutions:
  1615. attn.append(AttnBlock3D(block_in))
  1616. up = nn.Module()
  1617. up.block = block
  1618. up.attn = attn
  1619. if i_level != 0:
  1620. up.upsample = Upsample3D_HW(block_in, resamp_with_conv)
  1621. curr_h *= 2
  1622. curr_w *= 2
  1623. self.up.insert(0, up)
  1624. # end
  1625. self.norm_out = Normalize(block_in)
  1626. self.conv_out = nn.Conv3d(block_in, out_ch, kernel_size=3, stride=1, padding=1)
  1627. def forward(self, z):
  1628. temb = None
  1629. h = self.conv_in(z)
  1630. h = self.mid.block_1(h, temb)
  1631. h = self.mid.attn_1(h)
  1632. h = self.mid.block_2(h, temb)
  1633. for i_level in reversed(range(self.num_resolutions)):
  1634. for i_block in range(self.num_res_blocks + 1):
  1635. h = self.up[i_level].block[i_block](h, temb)
  1636. if len(self.up[i_level].attn) > 0:
  1637. h = self.up[i_level].attn[i_block](h)
  1638. if i_level != 0:
  1639. h = self.up[i_level].upsample(h)
  1640. if self.give_pre_end:
  1641. return h
  1642. h = self.norm_out(h)
  1643. h = nonlinearity(h)
  1644. h = self.conv_out(h)
  1645. return h
  1646. class Encoder3DFiLM(nn.Module):
  1647. """
  1648. 3D Encoder using ResnetBlock3DFiLM for voxel size conditioning.
  1649. Conditions on SOURCE voxel size.
  1650. """
  1651. def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
  1652. attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels,
  1653. resolution, z_channels, double_z=True,
  1654. cond_input_dim=3, cond_embed_dim=512, # FiLM parameters for SOURCE cond
  1655. **ignore_kwargs):
  1656. super().__init__()
  1657. self.ch = ch
  1658. self.temb_ch = 0 # No time embedding in VAE encoder typically
  1659. self.num_resolutions = len(ch_mult)
  1660. self.num_res_blocks = num_res_blocks
  1661. self.resolution = resolution
  1662. self.in_channels = in_channels
  1663. self.cond_input_dim = cond_input_dim
  1664. self.cond_embed_dim = cond_embed_dim
  1665. print(f"Encoder3DFiLM: Conditioning on input dim {self.cond_input_dim}, embedding to {self.cond_embed_dim}")
  1666. # Conditioning Embedder
  1667. if self.cond_embed_dim is not None and self.cond_embed_dim > 0:
  1668. self.cond_embedder = nn.Sequential(
  1669. nn.Linear(self.cond_input_dim, self.ch * 4),
  1670. nn.SiLU(),
  1671. nn.Linear(self.ch * 4, self.cond_embed_dim)
  1672. )
  1673. else:
  1674. print("Encoder3DFiLM: No conditioning embedding will be used.")
  1675. self.cond_embed_dim = None
  1676. self.cond_embedder = None
  1677. # downsampling
  1678. self.conv_in = torch.nn.Conv3d(in_channels,
  1679. self.ch,
  1680. kernel_size=3,
  1681. stride=1,
  1682. padding=1)
  1683. curr_res = resolution
  1684. in_ch_mult = (1,)+tuple(ch_mult)
  1685. self.down = nn.ModuleList()
  1686. for i_level in range(self.num_resolutions):
  1687. block = nn.ModuleList()
  1688. attn = nn.ModuleList()
  1689. block_in = ch*in_ch_mult[i_level]
  1690. block_out = ch*ch_mult[i_level]
  1691. for i_block in range(self.num_res_blocks):
  1692. block.append(ResnetBlock3DFiLM(in_channels=block_in, # Use FiLM block
  1693. out_channels=block_out,
  1694. temb_channels=self.temb_ch,
  1695. dropout=dropout,
  1696. cond_embed_dim=self.cond_embed_dim)) # Pass cond dim
  1697. block_in = block_out
  1698. # Calculate resolution at this level correctly BEFORE checking attn_resolutions
  1699. level_res = resolution // 2**(i_level)
  1700. if level_res in attn_resolutions:
  1701. attn.append(AttnBlock3D(block_in)) # Keep standard self-attention
  1702. down = nn.Module()
  1703. down.block = block
  1704. down.attn = attn # Store attention modules
  1705. if i_level != self.num_resolutions-1:
  1706. down.downsample = Downsample3D(block_in, resamp_with_conv)
  1707. # Update curr_res based on spatial dimensions if needed
  1708. curr_res = level_res // 2 # Assuming isotropic downsampling
  1709. self.down.append(down)
  1710. # Middle block
  1711. # block_in should be the output channels from the last downsampling level
  1712. block_in = ch*ch_mult[-1] # Correctly get channels for mid block
  1713. self.mid = nn.Module()
  1714. self.mid.block_1 = ResnetBlock3DFiLM(in_channels=block_in, # Use FiLM block
  1715. out_channels=block_in,
  1716. temb_channels=self.temb_ch,
  1717. dropout=dropout,
  1718. cond_embed_dim=self.cond_embed_dim)
  1719. self.mid.attn_1 = AttnBlock3D(block_in) # Standard self-attention
  1720. self.mid.block_2 = ResnetBlock3DFiLM(in_channels=block_in, # Use FiLM block
  1721. out_channels=block_in,
  1722. temb_channels=self.temb_ch,
  1723. dropout=dropout,
  1724. cond_embed_dim=self.cond_embed_dim)
  1725. # End block
  1726. self.norm_out = Normalize(block_in)
  1727. self.conv_out = torch.nn.Conv3d(block_in,
  1728. 2*z_channels if double_z else z_channels,
  1729. kernel_size=3,
  1730. stride=1,
  1731. padding=1)
  1732. def forward(self, x, cond_input=None):
  1733. # Generate conditioning embedding
  1734. cond_emb = None
  1735. if self.cond_embedder is not None and cond_input is not None:
  1736. # Ensure cond_input is a tensor
  1737. if not isinstance(cond_input, torch.Tensor):
  1738. # Create tensor explicitly on the SAME DEVICE as input x
  1739. cond_input = torch.tensor(cond_input, device=x.device, dtype=torch.float32)
  1740. # Ensure correct dtype
  1741. if cond_input.dtype != torch.float32:
  1742. cond_input = cond_input.to(dtype=torch.float32)
  1743. # <<< ADDED/MODIFIED: Ensure correct device >>>
  1744. if cond_input.device != x.device:
  1745. cond_input = cond_input.to(x.device)
  1746. # Handle batch dimension
  1747. if cond_input.ndim == 1:
  1748. cond_input = cond_input.unsqueeze(0).expand(x.shape[0], -1)
  1749. elif cond_input.shape[0] != x.shape[0]:
  1750. raise ValueError(f"Batch size mismatch: x {x.shape[0]}, cond {cond_input.shape[0]}")
  1751. cond_emb = self.cond_embedder(cond_input) # Now
  1752. temb = None # No time embedding
  1753. # Downsampling
  1754. hs = [self.conv_in(x)] # Store feature maps at each resolution
  1755. h = hs[-1] # Current feature map
  1756. for i_level in range(self.num_resolutions):
  1757. # Apply ResNet blocks for this level
  1758. for i_block in range(self.num_res_blocks):
  1759. res_block = self.down[i_level].block[i_block]
  1760. h = res_block(h, temb, cond_emb=cond_emb) # Update h in place
  1761. # Apply Attention blocks after ResNet blocks for this level
  1762. if len(self.down[i_level].attn) > 0:
  1763. for attn_block in self.down[i_level].attn:
  1764. h = attn_block(h) # Update h in place
  1765. hs.append(h) # Store the output of this level (before downsampling)
  1766. # Downsample if not the last level
  1767. if i_level != self.num_resolutions - 1:
  1768. h = self.down[i_level].downsample(h) # Downsample for next level's input
  1769. # We actually store the output *before* downsampling in hs for skip connections
  1770. # The input to the next level is the downsampled 'h'
  1771. # Middle block
  1772. # 'h' now holds the output of the last downsampling operation
  1773. h = self.mid.block_1(h, temb, cond_emb=cond_emb)
  1774. h = self.mid.attn_1(h)
  1775. h = self.mid.block_2(h, temb, cond_emb=cond_emb)
  1776. # End block
  1777. h = self.norm_out(h)
  1778. h = nonlinearity(h)
  1779. h = self.conv_out(h)
  1780. return h
  1781. class Decoder3DFiLM(nn.Module):
  1782. """
  1783. 3D Decoder using ResnetBlock3DFiLM for voxel size conditioning. (Corrected)
  1784. """
  1785. def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
  1786. attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels,
  1787. resolution, z_channels, give_pre_end=False,
  1788. cond_input_dim=3, cond_embed_dim=512, # FiLM parameters
  1789. **ignorekwargs):
  1790. super().__init__()
  1791. self.ch = ch
  1792. self.temb_ch = 0 # No time embedding in VAE decoder typically
  1793. self.num_resolutions = len(ch_mult)
  1794. self.num_res_blocks = num_res_blocks
  1795. self.resolution = resolution
  1796. self.in_channels = in_channels
  1797. self.give_pre_end = give_pre_end
  1798. self.cond_input_dim = cond_input_dim
  1799. self.cond_embed_dim = cond_embed_dim
  1800. in_ch_mult = (1,)+tuple(ch_mult)
  1801. block_in = ch*ch_mult[self.num_resolutions-1]
  1802. curr_res = resolution // 2**(self.num_resolutions-1)
  1803. # Adjust z_shape if resolution is not cubic
  1804. # Assuming resolution is isotropic for simplicity here
  1805. self.z_shape = (1, z_channels, curr_res, curr_res, curr_res)
  1806. print(f"Decoder3DFiLM: z_shape approx {self.z_shape}, dims {np.prod(self.z_shape)}")
  1807. print(f"Decoder3DFiLM: Conditioning on input dim {self.cond_input_dim}, embedding to {self.cond_embed_dim}")
  1808. # Conditioning Embedder
  1809. if self.cond_embed_dim is not None and self.cond_embed_dim > 0:
  1810. self.cond_embedder = nn.Sequential(
  1811. nn.Linear(self.cond_input_dim, self.ch * 4),
  1812. nn.SiLU(), # <<< CORRECT MODULE for Sequential
  1813. nn.Linear(self.ch * 4, self.cond_embed_dim)
  1814. )
  1815. else:
  1816. print("Decoder3DFiLM: No conditioning embedding will be used.")
  1817. self.cond_embed_dim = None
  1818. self.cond_embedder = None
  1819. # Input convolution
  1820. self.conv_in = nn.Conv3d(z_channels, block_in, kernel_size=3, stride=1, padding=1)
  1821. # Middle blocks
  1822. self.mid = nn.Module()
  1823. self.mid.block_1 = ResnetBlock3DFiLM(in_channels=block_in, out_channels=block_in,
  1824. temb_channels=self.temb_ch, dropout=dropout,
  1825. cond_embed_dim=self.cond_embed_dim)
  1826. self.mid.attn_1 = AttnBlock3D(block_in)
  1827. self.mid.block_2 = ResnetBlock3DFiLM(in_channels=block_in, out_channels=block_in,
  1828. temb_channels=self.temb_ch, dropout=dropout,
  1829. cond_embed_dim=self.cond_embed_dim)
  1830. # Upsampling blocks
  1831. self.up = nn.ModuleList()
  1832. for i_level in reversed(range(self.num_resolutions)):
  1833. block = nn.ModuleList()
  1834. attn = nn.ModuleList()
  1835. block_out = ch*ch_mult[i_level]
  1836. # Correct calculation of block_in for upsampling path
  1837. # It should be the output channels from the previous level (or mid block)
  1838. current_block_in = block_in # block_in holds channels from previous level
  1839. for i_block in range(self.num_res_blocks+1):
  1840. # ResNet block input channels = channels from previous layer in this level
  1841. block.append(ResnetBlock3DFiLM(in_channels=current_block_in, # Use correct input channels
  1842. out_channels=block_out,
  1843. temb_channels=self.temb_ch,
  1844. dropout=dropout,
  1845. cond_embed_dim=self.cond_embed_dim))
  1846. current_block_in = block_out # Update block_in for the next block in *this* level
  1847. # block_in now holds the final output channels for this level (block_out)
  1848. block_in = block_out # Update block_in for the *next* (lower index) level's input
  1849. up = nn.Module()
  1850. up.block = block
  1851. up.attn = nn.ModuleList() # Initialize attn list
  1852. # Calculate resolution at this level correctly BEFORE checking attn_resolutions
  1853. level_res = resolution // 2**(i_level)
  1854. if level_res in attn_resolutions:
  1855. up.attn.append(AttnBlock3D(block_in)) # Add attention if resolution matches
  1856. if i_level != 0:
  1857. # Upsample takes block_in (output channels of this level) as input channels
  1858. up.upsample = Upsample3D(block_in, resamp_with_conv)
  1859. # Update curr_res correctly (assuming isotropic for now)
  1860. curr_res = level_res * 2
  1861. self.up.insert(0, up) # Prepend
  1862. # Output layers
  1863. self.norm_out = Normalize(block_in) # norm_out uses the final block_in channels
  1864. self.conv_out = nn.Conv3d(block_in, out_ch, kernel_size=3, stride=1, padding=1)
  1865. def forward(self, z, cond_input=None):
  1866. # ... (other setup) ...
  1867. # Generate conditioning embedding
  1868. temb = None
  1869. cond_emb = None
  1870. if self.cond_embedder is not None and cond_input is not None:
  1871. # Ensure cond_input is a tensor
  1872. if not isinstance(cond_input, torch.Tensor):
  1873. # Create tensor explicitly on the SAME DEVICE as input z
  1874. cond_input = torch.tensor(cond_input, device=z.device, dtype=torch.float32)
  1875. # Ensure correct dtype
  1876. if cond_input.dtype != torch.float32:
  1877. cond_input = cond_input.to(dtype=torch.float32)
  1878. # <<< ADDED/MODIFIED: Ensure correct device >>>
  1879. if cond_input.device != z.device:
  1880. cond_input = cond_input.to(z.device)
  1881. # Handle batch dimension
  1882. if cond_input.ndim == 1:
  1883. cond_input = cond_input.unsqueeze(0).expand(z.shape[0], -1)
  1884. elif cond_input.shape[0] != z.shape[0]:
  1885. raise ValueError(f"Batch size mismatch: z {z.shape[0]}, cond {cond_input.shape[0]}")
  1886. cond_emb = self.cond_embedder(cond_input) # Now
  1887. # Initial convolution
  1888. h = self.conv_in(z)
  1889. # Middle blocks
  1890. h = self.mid.block_1(h, temb, cond_emb=cond_emb)
  1891. h = self.mid.attn_1(h)
  1892. h = self.mid.block_2(h, temb, cond_emb=cond_emb)
  1893. # Upsampling pathway
  1894. for i_level in reversed(range(self.num_resolutions)):
  1895. # Apply ResNet blocks for this level
  1896. for i_block in range(self.num_res_blocks+1):
  1897. h = self.up[i_level].block[i_block](h, temb, cond_emb=cond_emb)
  1898. # Apply Attention blocks for this level
  1899. if len(self.up[i_level].attn) > 0:
  1900. for attn_block in self.up[i_level].attn: # Iterate through attn blocks if multiple
  1901. h = attn_block(h)
  1902. # Upsample if not the last level
  1903. if i_level != 0:
  1904. h = self.up[i_level].upsample(h)
  1905. # Final output processing
  1906. if self.give_pre_end:
  1907. return h
  1908. h = self.norm_out(h)
  1909. h = nonlinearity(h) # <<< Using nonlinearity FUNCTION here is FINE
  1910. h = self.conv_out(h)
  1911. return h
  1912. class Encoder3D(nn.Module):
  1913. def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
  1914. attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels,
  1915. resolution, z_channels, double_z=True, resizing_pos=None, **ignore_kwargs):
  1916. super().__init__()
  1917. self.ch = ch
  1918. self.temb_ch = 0
  1919. self.num_resolutions = len(ch_mult)
  1920. self.num_res_blocks = num_res_blocks
  1921. self.resolution = resolution
  1922. self.in_channels = in_channels
  1923. self.resizing_pos = resizing_pos
  1924. if self.resizing_pos is None:
  1925. print("`resizing_pos` not found in ddconfig. Defaulting to downsampling at all but the last level.")
  1926. self.resizing_pos = [1] * (self.num_resolutions - 1) + [0]
  1927. # ───────────────────────────────────────────────────────────────────
  1928. # NEW: Validation to ensure list lengths match
  1929. # ───────────────────────────────────────────────────────────────────
  1930. if len(self.resizing_pos) != self.num_resolutions:
  1931. raise ValueError(
  1932. f"Configuration Error: The length of 'resizing_pos' ({len(self.resizing_pos)}) "
  1933. f"must be equal to the length of 'ch_mult' ({self.num_resolutions})."
  1934. )
  1935. # ───────────────────────────────────────────────────────────────────
  1936. # downsampling
  1937. self.conv_in = torch.nn.Conv3d(in_channels,
  1938. self.ch,
  1939. kernel_size=3,
  1940. stride=1,
  1941. padding=1)
  1942. curr_res = resolution
  1943. in_ch_mult = (1,)+tuple(ch_mult)
  1944. self.down = nn.ModuleList()
  1945. for i_level in range(self.num_resolutions):
  1946. block = nn.ModuleList()
  1947. attn = nn.ModuleList()
  1948. block_in = ch*in_ch_mult[i_level]
  1949. block_out = ch*ch_mult[i_level]
  1950. for i_block in range(self.num_res_blocks):
  1951. block.append(ResnetBlock3D(in_channels=block_in,
  1952. out_channels=block_out,
  1953. temb_channels=self.temb_ch,
  1954. dropout=dropout))
  1955. block_in = block_out
  1956. if curr_res in attn_resolutions:
  1957. attn.append(AttnBlock3D(block_in))
  1958. down = nn.Module()
  1959. down.block = block
  1960. down.attn = attn
  1961. if self.resizing_pos[i_level] == 1:
  1962. down.downsample = Downsample3D(block_in, resamp_with_conv)
  1963. curr_res = curr_res // 2
  1964. self.down.append(down)
  1965. # middle
  1966. self.mid = nn.Module()
  1967. self.mid.block_1 = ResnetBlock3D(in_channels=block_in, out_channels=block_in, temb_channels=self.temb_ch, dropout=dropout)
  1968. self.mid.attn_1 = AttnBlock3D(block_in)
  1969. self.mid.block_2 = ResnetBlock3D(in_channels=block_in, out_channels=block_in, temb_channels=self.temb_ch, dropout=dropout)
  1970. # end
  1971. self.norm_out = Normalize(block_in)
  1972. self.conv_out = torch.nn.Conv3d(block_in, 2*z_channels if double_z else z_channels, kernel_size=3, stride=1, padding=1)
  1973. def forward(self, x):
  1974. temb = None
  1975. hs = [self.conv_in(x)]
  1976. for i_level in range(self.num_resolutions):
  1977. for i_block in range(self.num_res_blocks):
  1978. h = self.down[i_level].block[i_block](hs[-1], temb)
  1979. if len(self.down[i_level].attn) > 0:
  1980. h = self.down[i_level].attn[i_block](h)
  1981. hs.append(h)
  1982. if hasattr(self.down[i_level], 'downsample'):
  1983. hs.append(self.down[i_level].downsample(hs[-1]))
  1984. h = hs[-1]
  1985. h = self.mid.block_1(h, temb)
  1986. h = self.mid.attn_1(h)
  1987. h = self.mid.block_2(h, temb)
  1988. h = self.norm_out(h)
  1989. h = nonlinearity(h)
  1990. h = self.conv_out(h)
  1991. return h
  1992. class Decoder3D(nn.Module):
  1993. def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
  1994. attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels,
  1995. resolution, z_channels, give_pre_end=False, resizing_pos=None, **ignorekwargs):
  1996. super().__init__()
  1997. self.ch = ch
  1998. self.temb_ch = 0
  1999. self.num_resolutions = len(ch_mult)
  2000. self.num_res_blocks = num_res_blocks
  2001. self.resolution = resolution
  2002. self.in_channels = in_channels
  2003. self.give_pre_end = give_pre_end
  2004. self.resizing_pos = resizing_pos
  2005. if self.resizing_pos is None:
  2006. print("`resizing_pos` not found in ddconfig. Defaulting to upsampling at all but the first level.")
  2007. self.resizing_pos = [1] * (self.num_resolutions - 1) + [0]
  2008. # ───────────────────────────────────────────────────────────────────
  2009. # NEW: Validation to ensure list lengths match
  2010. # ───────────────────────────────────────────────────────────────────
  2011. if len(self.resizing_pos) != self.num_resolutions:
  2012. raise ValueError(
  2013. f"Configuration Error: The length of 'resizing_pos' ({len(self.resizing_pos)}) "
  2014. f"must be equal to the length of 'ch_mult' ({self.num_resolutions})."
  2015. )
  2016. # ───────────────────────────────────────────────────────────────────
  2017. in_ch_mult = (1,)+tuple(ch_mult)
  2018. block_in = ch*ch_mult[self.num_resolutions-1]
  2019. num_downsamples = sum(self.resizing_pos)
  2020. curr_res = resolution // (2**num_downsamples)
  2021. self.z_shape = (1,z_channels,curr_res,curr_res)
  2022. print("Working with z of shape {} = {} dimensions.".format(self.z_shape, np.prod(self.z_shape)))
  2023. self.conv_in = torch.nn.Conv3d(z_channels, block_in, kernel_size=3, stride=1, padding=1)
  2024. self.mid = nn.Module()
  2025. self.mid.block_1 = ResnetBlock3D(in_channels=block_in, out_channels=block_in, temb_channels=self.temb_ch, dropout=dropout)
  2026. self.mid.attn_1 = AttnBlock3D(block_in)
  2027. self.mid.block_2 = ResnetBlock3D(in_channels=block_in, out_channels=block_in, temb_channels=self.temb_ch, dropout=dropout)
  2028. self.up = nn.ModuleList()
  2029. for i_level in reversed(range(self.num_resolutions)):
  2030. block = nn.ModuleList()
  2031. attn = nn.ModuleList()
  2032. block_out = ch*ch_mult[i_level]
  2033. for i_block in range(self.num_res_blocks+1):
  2034. block.append(ResnetBlock3D(in_channels=block_in,
  2035. out_channels=block_out,
  2036. temb_channels=self.temb_ch,
  2037. dropout=dropout))
  2038. block_in = block_out
  2039. if curr_res in attn_resolutions:
  2040. attn.append(AttnBlock3D(block_in))
  2041. up = nn.Module()
  2042. up.block = block
  2043. up.attn = attn
  2044. if i_level != 0 and self.resizing_pos[i_level-1] == 1:
  2045. up.upsample = Upsample3D(block_in, resamp_with_conv)
  2046. curr_res = curr_res * 2
  2047. self.up.insert(0, up)
  2048. self.norm_out = Normalize(block_in)
  2049. self.conv_out = torch.nn.Conv3d(block_in, out_ch, kernel_size=3, stride=1, padding=1)
  2050. def forward(self, z):
  2051. self.last_z_shape = z.shape
  2052. temb = None
  2053. h = self.conv_in(z)
  2054. h = self.mid.block_1(h, temb)
  2055. h = self.mid.attn_1(h)
  2056. h = self.mid.block_2(h, temb)
  2057. for i_level in reversed(range(self.num_resolutions)):
  2058. for i_block in range(self.num_res_blocks+1):
  2059. h = self.up[i_level].block[i_block](h, temb)
  2060. if len(self.up[i_level].attn) > 0:
  2061. h = self.up[i_level].attn[i_block](h)
  2062. if hasattr(self.up[i_level], 'upsample'):
  2063. h = self.up[i_level].upsample(h)
  2064. if self.give_pre_end:
  2065. return h
  2066. h = self.norm_out(h)
  2067. h = nonlinearity(h)
  2068. h = self.conv_out(h)
  2069. return h
  2070. # from flash_attn.flash_attn_interface import flash_attn_qkvpacked_func
  2071. # ---------------------------------------------------------------------
  2072. # Small helpers carried over from your code base
  2073. # ---------------------------------------------------------------------
  2074. def nonlinearity(x): # SiLU / Swish
  2075. return x * torch.sigmoid(x)
  2076. class Normalize(nn.GroupNorm): # identical to your original
  2077. def __init__(self, num_channels):
  2078. super().__init__(num_groups=32, num_channels=num_channels, eps=1e-6, affine=True)
  2079. # ---------------------------------------------------------------------
  2080. # MEMORY-FRIENDLY FLASH-ATTENTION IN 3-D
  2081. # ---------------------------------------------------------------------
  2082. from torch.nn.functional import scaled_dot_product_attention as sdpa
  2083. # from utils_flash import FlashCompatMixin
  2084. class Attention3D_v2(nn.Module, FlashCompatMixin):
  2085. """
  2086. [B,C,D,H,W] → Flash/efficient/math-SDP → residual, with
  2087. automatic num_heads adjustment to satisfy head_dim ≤ 128 ∧ 8 | head_dim.
  2088. """
  2089. def __init__(self, in_channels: int, num_heads: int = 8, dropout: float = 0.):
  2090. super().__init__()
  2091. # --- pick num_heads so head_dim is a multiple of 8 and ≤128 ----------
  2092. while (in_channels // num_heads) > 128 or (in_channels // num_heads) % 8:
  2093. num_heads *= 2 # fall back to more heads
  2094. self.num_heads, self.head_dim = num_heads, in_channels // num_heads
  2095. if in_channels % num_heads:
  2096. raise ValueError(f"{in_channels=} not divisible by {num_heads=}")
  2097. inner = in_channels
  2098. self.norm = Normalize(in_channels)
  2099. self.qkv = nn.Conv3d(in_channels, inner * 3, 1)
  2100. self.proj_out = nn.Conv3d(inner, in_channels, 1)
  2101. self.dropout_p = dropout
  2102. # Inform user once
  2103. if not hasattr(Attention3D_v2, "_banner"):
  2104. dtype = torch.float16 # typical; real dtype known only after to(device)
  2105. self._explain_flash(self.head_dim, self.num_heads, dtype)
  2106. Attention3D_v2._banner = True
  2107. # ------------------------------------------------------------------
  2108. def forward(self, x):
  2109. B, C, D, H, W = x.shape
  2110. qkv = self.qkv(self.norm(x)).reshape(
  2111. B, 3, self.num_heads, self.head_dim, -1) # -1 = S=D·H·W
  2112. q, k, v = (t.permute(0,1,3,2) for t in qkv) # B,h,S,d
  2113. attn = sdpa(
  2114. q, k, v,
  2115. dropout_p=self.dropout_p if self.training else 0.,
  2116. is_causal=False
  2117. ) # B,h,S,d
  2118. attn = attn.permute(0,1,3,2).contiguous().view(B, C, D, H, W)
  2119. return x + self.proj_out(attn) # residual
  2120. # ---------------------------------------------------------------------
  2121. # BUILDING BLOCKS (unchanged logic, only suffixed)
  2122. # ---------------------------------------------------------------------
  2123. class ResnetBlock3D_v2(nn.Module):
  2124. def __init__(self, *, in_channels, out_channels=None, conv_shortcut=False,
  2125. dropout, temb_channels=512):
  2126. super().__init__()
  2127. out_channels = in_channels if out_channels is None else out_channels
  2128. self.in_channels, self.out_channels = in_channels, out_channels
  2129. self.use_conv_shortcut = conv_shortcut
  2130. self.norm1, self.norm2 = Normalize(in_channels), Normalize(out_channels)
  2131. self.conv1 = nn.Conv3d(in_channels, out_channels, 3, padding=1)
  2132. self.conv2 = nn.Conv3d(out_channels, out_channels, 3, padding=1)
  2133. self.dropout = nn.Dropout(dropout)
  2134. if temb_channels > 0:
  2135. self.temb_proj = nn.Linear(temb_channels, out_channels)
  2136. # skip connection if in!=out
  2137. if in_channels != out_channels:
  2138. self.conv_shortcut = (nn.Conv3d(in_channels, out_channels, 3, padding=1)
  2139. if conv_shortcut
  2140. else nn.Conv3d(in_channels, out_channels, 1))
  2141. def forward(self, x, temb):
  2142. h = self.conv1(nonlinearity(self.norm1(x)))
  2143. if temb is not None:
  2144. h += self.temb_proj(nonlinearity(temb)).unsqueeze(-1).unsqueeze(-1).unsqueeze(-1)
  2145. h = self.conv2(self.dropout(nonlinearity(self.norm2(h))))
  2146. if self.in_channels != self.out_channels:
  2147. x = self.conv_shortcut(x)
  2148. return x + h
  2149. class Upsample3D_v2(nn.Module):
  2150. def __init__(self, in_channels, with_conv):
  2151. super().__init__()
  2152. self.with_conv = with_conv
  2153. if with_conv:
  2154. self.conv = nn.Conv3d(in_channels, in_channels, 3, padding=1)
  2155. def forward(self, x):
  2156. x = F.interpolate(x, scale_factor=(2,2,2), mode="nearest")
  2157. return self.conv(x) if self.with_conv else x
  2158. class Downsample3D_v2(nn.Module):
  2159. def __init__(self, in_channels, with_conv):
  2160. super().__init__()
  2161. self.with_conv = with_conv
  2162. if with_conv:
  2163. self.conv = nn.Conv3d(in_channels, in_channels, 3, stride=2)
  2164. def forward(self, x):
  2165. if self.with_conv:
  2166. # pad so D,H,W divisible by 2
  2167. x = F.pad(x, (0,1,0,1,0,1))
  2168. x = self.conv(x)
  2169. else:
  2170. x = F.avg_pool3d(x, 2, stride=2)
  2171. return x
  2172. class Normalize(nn.Module):
  2173. """
  2174. A normalization layer, typically GroupNorm.
  2175. It robustly adjusts num_groups to be a divisor of num_channels.
  2176. """
  2177. def __init__(self, num_channels, num_groups=32):
  2178. super().__init__()
  2179. # Input validation for num_channels and num_groups
  2180. if not isinstance(num_channels, int) or num_channels < 0:
  2181. raise ValueError("num_channels must be a non-negative integer.")
  2182. if not isinstance(num_groups, int) or num_groups <= 0:
  2183. raise ValueError("num_groups must be a positive integer.")
  2184. if num_channels == 0: # Edge case: no normalization needed or possible for 0 channels
  2185. self.norm = nn.Identity()
  2186. elif num_channels < num_groups or num_channels % num_groups != 0:
  2187. # If num_channels is small (e.g., less than num_groups),
  2188. # or not divisible by num_groups, find a suitable num_groups.
  2189. if num_channels <= num_groups:
  2190. # Use num_channels as num_groups if it's smaller or equal.
  2191. # This means each channel is its own group.
  2192. num_groups = num_channels
  2193. else:
  2194. # Find the largest valid number of groups <= preferred num_groups
  2195. # that divides num_channels.
  2196. possible_num_groups = [g for g in range(1, min(num_groups, num_channels) + 1) if num_channels % g == 0]
  2197. if not possible_num_groups: # Should not happen if g=1 is always possible
  2198. num_groups = 1 # Fallback: normalize over all channels as a single group
  2199. else:
  2200. num_groups = max(possible_num_groups)
  2201. self.norm = nn.GroupNorm(num_groups=num_groups, num_channels=num_channels, eps=1e-6, affine=True)
  2202. else:
  2203. # num_channels is divisible by num_groups
  2204. self.norm = nn.GroupNorm(num_groups=num_groups, num_channels=num_channels, eps=1e-6, affine=True)
  2205. def forward(self, x):
  2206. # Skip normalization if input channels is 0 (e.g. empty tensor placeholder)
  2207. if x.shape[1] == 0:
  2208. return x
  2209. return self.norm(x)
  2210. class CrossResolutionAttention3D(nn.Module):
  2211. """
  2212. Attention mechanism where the original 3D input (queries) attends to
  2213. a downsampled version of itself (keys and values), using
  2214. torch.nn.functional.scaled_dot_product_attention.
  2215. """
  2216. def __init__(self, in_channels, pool_type='max', downsample_factor=2, num_norm_groups=32, attention_dropout_p=0.0):
  2217. super().__init__()
  2218. # Input validation
  2219. if not isinstance(in_channels, int) or in_channels <= 0:
  2220. raise ValueError("in_channels must be a positive integer.")
  2221. if not isinstance(downsample_factor, int) or downsample_factor < 1:
  2222. raise ValueError("downsample_factor must be an integer >= 1.")
  2223. if pool_type not in ['max', 'avg']:
  2224. raise ValueError(f"Unsupported pool_type: {pool_type}. Choose 'max' or 'avg'.")
  2225. if not (0.0 <= attention_dropout_p < 1.0):
  2226. raise ValueError("attention_dropout_p must be between 0.0 (inclusive) and 1.0 (exclusive).")
  2227. self.in_channels = in_channels
  2228. self.downsample_factor = downsample_factor
  2229. self.attention_dropout_p = attention_dropout_p
  2230. self.norm = Normalize(in_channels, num_groups=num_norm_groups)
  2231. # Projections for Q from original resolution feature map
  2232. self.q_conv = nn.Conv3d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
  2233. # Pooling layer
  2234. if self.downsample_factor > 1:
  2235. if pool_type == 'max':
  2236. self.pool = nn.MaxPool3d(kernel_size=downsample_factor, stride=downsample_factor)
  2237. elif pool_type == 'avg': # pool_type == 'avg'
  2238. self.pool = nn.AvgPool3d(kernel_size=downsample_factor, stride=downsample_factor)
  2239. else: # downsample_factor is 1, no actual pooling needed
  2240. self.pool = nn.Identity()
  2241. # Projections for K and V from the downsampled feature map
  2242. self.k_conv = nn.Conv3d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
  2243. self.v_conv = nn.Conv3d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
  2244. # Output projection
  2245. self.proj_out = nn.Conv3d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
  2246. def forward(self, x):
  2247. # b: batch_size, c: channels, h,w,z: spatial dimensions (e.g., Depth, Height, Width)
  2248. b, c, h, w, z = x.shape
  2249. if c != self.in_channels:
  2250. raise ValueError(f"Input channels {c} does not match model's initialized in_channels {self.in_channels}")
  2251. # 1. Normalize original input
  2252. x_norm = self.norm(x)
  2253. # 2. Generate Queries (Q) from original resolution
  2254. # q_spatial shape: (b, c, h, w, z)
  2255. q_spatial = self.q_conv(x_norm)
  2256. num_orig_voxels = h * w * z
  2257. # Reshape Q for attention: (b, N_orig, c) where N_orig = h*w*z
  2258. # (b, c, N_orig) -> (b, N_orig, c)
  2259. query = q_spatial.view(b, c, num_orig_voxels).permute(0, 2, 1)
  2260. # 3. Create downsampled version for Keys (K) and Values (V)
  2261. x_downsampled_norm = self.pool(x_norm) # Applies pooling if downsample_factor > 1
  2262. # k_spatial_downsampled shape: (b, c, h_d, w_d, z_d)
  2263. k_spatial_downsampled = self.k_conv(x_downsampled_norm)
  2264. # v_spatial_downsampled shape: (b, c, h_d, w_d, z_d)
  2265. v_spatial_downsampled = self.v_conv(x_downsampled_norm)
  2266. _, _, h_d, w_d, z_d = k_spatial_downsampled.shape # Dimensions of downsampled map
  2267. num_down_voxels = h_d * w_d * z_d
  2268. # Reshape K for attention: (b, N_down, c)
  2269. # (b, c, N_down) -> (b, N_down, c)
  2270. key = k_spatial_downsampled.view(b, c, num_down_voxels).permute(0, 2, 1)
  2271. # Reshape V for attention: (b, N_down, c)
  2272. # (b, c, N_down) -> (b, N_down, c)
  2273. value = v_spatial_downsampled.view(b, c, num_down_voxels).permute(0, 2, 1)
  2274. # 4. Compute attention using scaled_dot_product_attention
  2275. # query shape: (b, num_orig_voxels, c)
  2276. # key shape: (b, num_down_voxels, c)
  2277. # value shape: (b, num_down_voxels, c)
  2278. # Output shape: (b, num_orig_voxels, c)
  2279. attn_output_reshaped = F.scaled_dot_product_attention(
  2280. query,
  2281. key,
  2282. value,
  2283. attn_mask=None, # No mask typically needed for this non-causal cross-attention
  2284. dropout_p=self.attention_dropout_p if self.training else 0.0, # Apply dropout only during training
  2285. is_causal=False
  2286. )
  2287. # 5. Reshape attended output back to original spatial dimensions
  2288. # (b, N_orig, c) -> (b, c, N_orig) -> (b, c, h, w, z)
  2289. attn_output_spatial = attn_output_reshaped.permute(0, 2, 1).contiguous().view(b, c, h, w, z)
  2290. # 6. Final projection
  2291. projected_output = self.proj_out(attn_output_spatial)
  2292. # 7. Residual connection
  2293. return x + projected_output
  2294. # ---------------------------------------------------------------------
  2295. # ENCODER / DECODER
  2296. # ---------------------------------------------------------------------
  2297. class Encoder3D_v3(nn.Module):
  2298. def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
  2299. attn_resolutions, dropout=0., resamp_with_conv=True, in_channels,
  2300. resolution, z_channels, double_z=True, resizing_pos=None, **ignore):
  2301. super().__init__()
  2302. self.ch, self.temb_ch = ch, 0
  2303. n_levels = len(ch_mult)
  2304. self.num_res_blocks, self.resolution, self.in_channels = num_res_blocks, resolution, in_channels
  2305. self.resizing_pos = resizing_pos or ([1]*(n_levels-1) + [0])
  2306. assert len(self.resizing_pos) == n_levels, "`resizing_pos` length mismatch."
  2307. self.conv_in = nn.Conv3d(in_channels, ch, 3, padding=1)
  2308. curr_res, in_ch_mult = resolution, (1,)+tuple(ch_mult)
  2309. self.down = nn.ModuleList()
  2310. for lvl in range(n_levels):
  2311. block, attn = nn.ModuleList(), nn.ModuleList()
  2312. in_ch = ch * in_ch_mult[lvl]
  2313. out_ch = ch * ch_mult[lvl]
  2314. for _ in range(num_res_blocks):
  2315. block.append(ResnetBlock3D_v2(in_channels=in_ch, out_channels=out_ch,
  2316. temb_channels=self.temb_ch, dropout=dropout))
  2317. in_ch = out_ch
  2318. if curr_res in attn_resolutions:
  2319. attn.append(CrossResolutionAttention3D(in_ch))
  2320. down_l = nn.Module(); down_l.block, down_l.attn = block, attn
  2321. if self.resizing_pos[lvl]:
  2322. down_l.downsample = Downsample3D_v2(in_ch, resamp_with_conv)
  2323. curr_res //= 2
  2324. self.down.append(down_l)
  2325. # middle
  2326. self.mid = nn.Module()
  2327. self.mid.block_1 = ResnetBlock3D_v2(in_channels=in_ch, dropout=dropout)
  2328. self.mid.attn_1 = CrossResolutionAttention3D(in_ch)
  2329. self.mid.block_2 = ResnetBlock3D_v2(in_channels=in_ch, dropout=dropout)
  2330. # output
  2331. self.norm_out = Normalize(in_ch)
  2332. self.conv_out = nn.Conv3d(in_ch, 2*z_channels if double_z else z_channels, 3, padding=1)
  2333. def forward(self, x):
  2334. temb, hs = None, [self.conv_in(x)]
  2335. for lvl in range(len(self.down)):
  2336. for blk in self.down[lvl].block:
  2337. h = blk(hs[-1], temb); hs.append(h)
  2338. if self.down[lvl].attn: h = self.down[lvl].attn[len(hs)-2](h); hs[-1] = h
  2339. if hasattr(self.down[lvl], "downsample"):
  2340. hs.append(self.down[lvl].downsample(hs[-1]))
  2341. h = hs[-1]; h = self.mid.block_1(h, temb); h = self.mid.attn_1(h); h = self.mid.block_2(h, temb)
  2342. return self.conv_out(nonlinearity(self.norm_out(h)))
  2343. class Decoder3D_v3(nn.Module):
  2344. def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks,
  2345. attn_resolutions, dropout=0., resamp_with_conv=True, in_channels,
  2346. resolution, z_channels, give_pre_end=False, resizing_pos=None, **kw):
  2347. super().__init__()
  2348. self.ch, self.temb_ch = ch, 0
  2349. n_levels = len(ch_mult)
  2350. self.num_res_blocks, self.give_pre_end = num_res_blocks, give_pre_end
  2351. self.resizing_pos = resizing_pos or ([1]*(n_levels-1)+[0])
  2352. assert len(self.resizing_pos) == n_levels
  2353. block_in = ch * ch_mult[-1]
  2354. num_down = sum(self.resizing_pos)
  2355. curr_res = resolution // (2 ** num_down)
  2356. self.z_shape = (1, z_channels, curr_res, curr_res, curr_res)
  2357. self.conv_in = nn.Conv3d(z_channels, block_in, 3, padding=1)
  2358. self.mid = nn.Module()
  2359. self.mid.block_1 = ResnetBlock3D_v2(in_channels=block_in, dropout=dropout)
  2360. self.mid.attn_1 = CrossResolutionAttention3D(block_in)
  2361. self.mid.block_2 = ResnetBlock3D_v2(in_channels=block_in, dropout=dropout)
  2362. self.up = nn.ModuleList()
  2363. for lvl in reversed(range(n_levels)):
  2364. block, attn = nn.ModuleList(), nn.ModuleList()
  2365. block_out = ch * ch_mult[lvl]
  2366. for _ in range(num_res_blocks + 1):
  2367. block.append(ResnetBlock3D_v2(in_channels=block_in, out_channels=block_out,
  2368. temb_channels=self.temb_ch, dropout=dropout))
  2369. block_in = block_out
  2370. if curr_res in attn_resolutions:
  2371. attn.append(CrossResolutionAttention3D(block_in))
  2372. up_l = nn.Module(); up_l.block, up_l.attn = block, attn
  2373. if lvl != 0 and self.resizing_pos[lvl-1]:
  2374. up_l.upsample = Upsample3D_v2(block_in, resamp_with_conv)
  2375. curr_res *= 2
  2376. self.up.insert(0, up_l)
  2377. self.norm_out = Normalize(block_in)
  2378. self.conv_out = nn.Conv3d(block_in, out_ch, 3, padding=1)
  2379. def forward(self, z):
  2380. temb, h = None, self.conv_in(z)
  2381. h = self.mid.block_1(h, temb); h = self.mid.attn_1(h); h = self.mid.block_2(h, temb)
  2382. for lvl in reversed(range(len(self.up))):
  2383. for blk in self.up[lvl].block:
  2384. h = blk(h, temb)
  2385. if self.up[lvl].attn: h = self.up[lvl].attn[0](h)
  2386. if hasattr(self.up[lvl], "upsample"):
  2387. h = self.up[lvl].upsample(h)
  2388. if self.give_pre_end:
  2389. return h
  2390. return self.conv_out(nonlinearity(self.norm_out(h)))

model.py at commit a16c2fe, under MIT · at the source

Overview

  1. Integrated Program in Neuroscience, McGill University, Montreal, Quebec, Canada
  2. Douglas Mental Health University Institute, Verdun, Quebec, Canada
  3. Department of Medicine, Université de Montréal, Montreal, Quebec, Canada
  4. Department of Psychiatry, McGill University, Montreal, Quebec, Canada
Journal: Imaging neuroscience (Cambridge, Mass.), volume 4, article IMAG.a.1326
Dates: received 28 September 2025; accepted 23 June 2026; published online 4 August 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1162/imag.a.1326 · PMID 42559373 · PMCID PMC13440154 · OpenAlex W7168404105
Open access: diamond, a free copy (OpenAlex)
Status: code verified
Categories: structural MRI / diffusion (modality)
Methods: Connectivity, Statistics, Machine learning
Keywords: single image super-resolution, magnetic resonance imaging, deep learning, contrast agnostic, resolution agnostic
Topic: Generative Adversarial Networks and Image Synthesis (Computer Vision and Pattern Recognition, Computer Science), according to OpenAlex
Funding: Healthy Brains for Healthy Lives (HBHL); Alzheimer Society Research Program (ASRP); Douglas Research Centre (DRC); Natural Sciences and Engineering Research Council of Canada (NSERC); Canadian Institutes of Health Research (CIHR); Fonds de Recherche du Quebec—Santé (FRQS)
Citations: not cited yet (Europe PMC); 58 references in the paper

Abstract

Due to their high inter-tissue contrast, Magnetic resonance images (MRIs) can reflect neuroanatomical changes related to healthy aging and pathological processes. However, standard brain MRI acquisition resolutions hinder the ability to measure the more subtle changes that occur in early disease stages. Increasing the resolution during acquisition poses multiple challenges, including increased noise, higher acquisition times and cost, and discomfort of the scanned individual. In this work, we propose a robust, generalizable single-image super-resolution network for brain MRIs named Resolution Augmentation with Variational auto-Encoder Networks (RAVEN) with generative adversarial networks (GANs). We show RAVEN is capable of upsampling in-vivo and ex-vivo MRIs of diverse modalities (e.g. T1-weighted, T2-weighted, and T2*) and varying field strengths (3T to 7T) to target voxel sizes as small as 0.5 mm isotropic using arbitrary upsampling factors. RAVEN achieved state-of-the-art performance against deep learning and non-deep learning methods, best preserving true anatomical information. We have also made RAVEN open access, with the source code as well as training and evaluation scripts available and ready to use at: https://github.com/waadgo/raven.

Reproduced under the paper's license (CC BY), from the paper cited above.

Repositories

Its files are read in the Code ↔ Paper reader above, with 9 matches between paragraphs and lines of code.

richzhang/PerceptualSimilarity

License: BSD-2-Clause
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 082bb24f84c091ea94de2867d34c4544f68e0963, 19 December 2023
Languages: Python (25), Shell (6)
Size: 55 files, 31 scripts
Software Heritage: not archived
Found in: the text, “Synthetic benchmark”
Holds: README, license file, environment (Dockerfile, requirements.txt, setup.py)
Not found: CITATION.cff, tests, continuous integration, documentation
Tools: PyTorch (14 files), NumPy (12 files), Matplotlib (5 files), Pillow (4 files), SciPy (3 files), OpenCV (2 files), scikit-image (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
33 files

waadgo/raven

License: MIT
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: a16c2fe61b91871694f0ef759a657780df8f5375, 18 August 2026
Languages: Python (320), Shell (8)
Size: 494 files, 328 scripts
Software Heritage: not archived
Found in: “Data and Code Availability”
Holds: README, license file, environment (requirements-raven.txt, setup.py, inference_runtime/setup.py, provenance/source_environment/environment.yaml, provenance/source_environment/environment_taming.yaml, provenance/source_environment/requirements.txt, provenance/source_environment/requirements_final.txt, provenance/source_environment/requirements_taming.txt, inference/raven_v1_3/runtime/setup.py), tests, documentation
Not found: CITATION.cff, continuous integration
Tools: PyTorch (190 files), NumPy (155 files), Pillow (45 files), h5py (44 files), Matplotlib (41 files), pandas (34 files), NiBabel (33 files), PyTorch Lightning (26 files), SciPy (20 files), OpenCV (18 files), SimpleITK (10 files), scikit-image (8 files), Hugging Face Transformers (6 files), imageio (3 files)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
330 files

The paper's code and data availability statement is in the Data section.

Tracing map

Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.

What the map holds:

  • 2 repositories of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 359 scripts, each with its path and the digest of its content;
  • 9 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.

Data

Datasets cited

Data and Code Availability

All the datasets used in this study (except for DBCBB) are available at the following URLs:

ADNI: https://adni.loni.usc.edu/data-samples/adni-data/neuroimaging/mri/mri-image-data-sets/; AHEAD: https://uvaauas.figshare.com/articles/dataset/The_Amsterdam_Ultra-high_field_adult_lifespan_database_AHEAD_A_freely_available_multimodal_7_Tesla_submillimeter_magnetic_resonance_imaging_database/10007840; AOMIC: https://nilab-uva.github.io/AOMIC.github.io/; CHEN: https://springernature.figshare.com/articles/dataset/UNC_Paired_3T-7T_Dataset/23706033; HBA: https://osf.io/ckh5t/files/osfstorage; HCP: https://www.humanconnectome.org/; OSCBS: https://openscience.cbs.mpg.de/bazin/7T_Quantitative/; LUSEBRINK: https://openneuro.org/datasets/ds003563/versions/1.1.0;

The code and trained weights are publicly available on https://github.com/waadgo/raven.

Reproduced under the paper's license (CC BY), from the paper cited above.

Versions

The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, pages, dates, 4 authors, 5 keywords, 6 funders, 52 references.

Cite

This paper

Adame-Gonzalez, W., Moqadam, R., Zeighami, Y., & Dadar, M. (2026). RAVEN: Robust, generalizable, multi-resolution structural MRI upsampling using autoencoders. Imaging neuroscience (Cambridge, Mass.), 4, IMAG.a.1326. https://doi.org/10.1162/imag.a.1326

BibTeX

@article{adamegonzalez2026raven,
author = {Adame-Gonzalez, Walter and Moqadam, Roqaie and Zeighami, Yashar and Dadar, Mahsa},
title = {{RAVEN: Robust, generalizable, multi-resolution structural MRI upsampling using autoencoders}},
journal = {Imaging neuroscience (Cambridge, Mass.)},
year = {2026},
month = aug,
volume = {4},
pages = {IMAG.a.1326},
publisher = {MIT Press},
issn = {2837-6056},
doi = {10.1162/imag.a.1326},
url = {https://doi.org/10.1162/imag.a.1326},
pmid = {42559373},
pmcid = {PMC13440154}
}

RIS

TY - JOUR
AU - Adame-Gonzalez, Walter
AU - Moqadam, Roqaie
AU - Zeighami, Yashar
AU - Dadar, Mahsa
TI - RAVEN: Robust, generalizable, multi-resolution structural MRI upsampling using autoencoders
T2 - Imaging neuroscience (Cambridge, Mass.)
J2 - Imaging Neurosci (Camb)
PY - 2026
DA - 2026/08/04
VL - 4
SP - IMAG.a.1326
SN - 2837-6056
PB - MIT Press
DO - 10.1162/imag.a.1326
UR - https://doi.org/10.1162/imag.a.1326
LA - en
ER -

CSL-JSON

{
"id": "10.1162/imag.a.1326",
"type": "article-journal",
"title": "RAVEN: Robust, generalizable, multi-resolution structural MRI upsampling using autoencoders",
"container-title": "Imaging neuroscience (Cambridge, Mass.)",
"author": [
{
"family": "Adame-Gonzalez",
"given": "Walter"
},
{
"family": "Moqadam",
"given": "Roqaie"
},
{
"family": "Zeighami",
"given": "Yashar"
},
{
"family": "Dadar",
"given": "Mahsa"
}
],
"container-title-short": "Imaging Neurosci (Camb)",
"volume": "4",
"page": "IMAG.a.1326",
"DOI": "10.1162/imag.a.1326",
"PMID": "42559373",
"PMCID": "PMC13440154",
"ISSN": "2837-6056",
"publisher": "MIT Press",
"URL": "https://doi.org/10.1162/imag.a.1326",
"language": "en",
"issued": {
"date-parts": [
[
2026,
8,
4
]
]
}
}

The tracing map gets a citation of its own once an author has validated it and it has a DOI.

Similar papers

The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.

[1] doi:10.1162/imag.a.1235 [code]
Intracranial volume: To adjust or not to adjust? It is not a matter of if, but how.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: pandas, SciPy, NumPy, structural MRI / diffusion, 1 reference, 3 authors
[2] doi:10.1371/journal.pcbi.1014263 [code]
MIRAGE: Robust multi-modal architectures translate fMRI-to-image models from vision to mental imagery.
Journal: PLoS computational biology
In common: PyTorch Lightning, imageio, Hugging Face Transformers, 10 other tools
[3] doi:10.3389/frai.2026.1771088 [code]
Few-shot deployment of pretrained MRI transformers in brain imaging tasks.
Journal: Frontiers in artificial intelligence
In common: imageio, SimpleITK, OpenCV, 9 other tools, structural MRI / diffusion
[4] doi:10.1038/s41467-026-73373-w [code]
Mapping neuro-vascular unit communications reveals distinct angiogenic programs across developing mouse brain regions.
Journal: Nature communications
In common: imageio, SimpleITK, OpenCV, 9 other tools
[5] doi:10.1364/boe.605322 [code]
Generalized plaque digitization framework for multi-dimensional mesoscopic images.
Journal: Biomedical optics express
In common: imageio, SimpleITK, OpenCV, 8 other tools
[6] doi:10.1111/joa.70203 [code]
Two-step workflow integrating automatic registration and manual refinement for the accurate alignment of serial histological sections in 3D reconstruction.
Journal: Journal of anatomy
In common: Hugging Face Transformers, SimpleITK, OpenCV, 8 other tools
[7] doi:10.1038/s41467-026-71555-0 [code]
A deep representation learning model to predict response to vagus nerve stimulation.
Journal: Nature communications
In common: PyTorch Lightning, NiBabel, PyTorch, 4 other tools, structural MRI / diffusion, 4 references
[8] doi:10.1002/alz.71649 [code]
Postmortem brain MRI reveals differential associations of subcortical and limbic volumes with cortical thinning and neurodegenerative pathologies.
Journal: Alzheimer's & dementia : the journal of the Alzheimer's Association
In common: SimpleITK, OpenCV, scikit-image, 7 other tools, structural MRI / diffusion, 1 reference
[9] doi:10.1038/s41598-026-57519-w [code]
Automated segmentation of neurons and spinal cord structures in immunofluorescence images using SpineDL.
Journal: Scientific reports
In common: imageio, OpenCV, scikit-image, 8 other tools
[10] doi:10.1038/s41597-026-07248-6 [code]
A large-scale fMRI dataset for vision-language semantic association.
Journal: Scientific data
In common: imageio, OpenCV, scikit-image, 8 other tools

Contribute

The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.

Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.

Request its removal

To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).

Discussion, reproductions, activity

Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.

Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.

Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.