candle_transformers/models/stable_diffusion/
resnet.rs1use crate::models::with_tracing::{conv2d, Conv2d};
9use candle::{Result, Tensor, D};
10use candle_nn as nn;
11use candle_nn::Module;
12
13#[derive(Debug, Clone, Copy)]
15pub struct ResnetBlock2DConfig {
16 pub out_channels: Option<usize>,
18 pub temb_channels: Option<usize>,
19 pub groups: usize,
21 pub groups_out: Option<usize>,
22 pub eps: f64,
24 pub use_in_shortcut: Option<bool>,
28 pub output_scale_factor: f64,
31}
32
33impl Default for ResnetBlock2DConfig {
34 fn default() -> Self {
35 Self {
36 out_channels: None,
37 temb_channels: Some(512),
38 groups: 32,
39 groups_out: None,
40 eps: 1e-6,
41 use_in_shortcut: None,
42 output_scale_factor: 1.,
43 }
44 }
45}
46
47#[derive(Debug)]
48pub struct ResnetBlock2D {
49 norm1: nn::GroupNorm,
50 conv1: Conv2d,
51 norm2: nn::GroupNorm,
52 conv2: Conv2d,
53 time_emb_proj: Option<nn::Linear>,
54 conv_shortcut: Option<Conv2d>,
55 span: tracing::Span,
56 config: ResnetBlock2DConfig,
57}
58
59impl ResnetBlock2D {
60 pub fn new(
61 vs: nn::VarBuilder,
62 in_channels: usize,
63 config: ResnetBlock2DConfig,
64 ) -> Result<Self> {
65 let out_channels = config.out_channels.unwrap_or(in_channels);
66 let conv_cfg = nn::Conv2dConfig {
67 stride: 1,
68 padding: 1,
69 groups: 1,
70 dilation: 1,
71 cudnn_fwd_algo: None,
72 };
73 let norm1 = nn::group_norm(config.groups, in_channels, config.eps, vs.pp("norm1"))?;
74 let conv1 = conv2d(in_channels, out_channels, 3, conv_cfg, vs.pp("conv1"))?;
75 let groups_out = config.groups_out.unwrap_or(config.groups);
76 let norm2 = nn::group_norm(groups_out, out_channels, config.eps, vs.pp("norm2"))?;
77 let conv2 = conv2d(out_channels, out_channels, 3, conv_cfg, vs.pp("conv2"))?;
78 let use_in_shortcut = config
79 .use_in_shortcut
80 .unwrap_or(in_channels != out_channels);
81 let conv_shortcut = if use_in_shortcut {
82 let conv_cfg = nn::Conv2dConfig {
83 stride: 1,
84 padding: 0,
85 groups: 1,
86 dilation: 1,
87 cudnn_fwd_algo: None,
88 };
89 Some(conv2d(
90 in_channels,
91 out_channels,
92 1,
93 conv_cfg,
94 vs.pp("conv_shortcut"),
95 )?)
96 } else {
97 None
98 };
99 let time_emb_proj = match config.temb_channels {
100 None => None,
101 Some(temb_channels) => Some(nn::linear(
102 temb_channels,
103 out_channels,
104 vs.pp("time_emb_proj"),
105 )?),
106 };
107 let span = tracing::span!(tracing::Level::TRACE, "resnet2d");
108 Ok(Self {
109 norm1,
110 conv1,
111 norm2,
112 conv2,
113 time_emb_proj,
114 span,
115 config,
116 conv_shortcut,
117 })
118 }
119
120 pub fn forward(&self, xs: &Tensor, temb: Option<&Tensor>) -> Result<Tensor> {
121 let _enter = self.span.enter();
122 let shortcut_xs = match &self.conv_shortcut {
123 Some(conv_shortcut) => conv_shortcut.forward(xs)?,
124 None => xs.clone(),
125 };
126 let xs = self.norm1.forward(xs)?;
127 let xs = self.conv1.forward(&nn::ops::silu(&xs)?)?;
128 let xs = match (temb, &self.time_emb_proj) {
129 (Some(temb), Some(time_emb_proj)) => time_emb_proj
130 .forward(&nn::ops::silu(temb)?)?
131 .unsqueeze(D::Minus1)?
132 .unsqueeze(D::Minus1)?
133 .broadcast_add(&xs)?,
134 _ => xs,
135 };
136 let xs = self
137 .conv2
138 .forward(&nn::ops::silu(&self.norm2.forward(&xs)?)?)?;
139 (shortcut_xs + xs)? / self.config.output_scale_factor
140 }
141}