use more comfy ops
This commit is contained in:
@@ -27,6 +27,9 @@ from functools import partial
|
||||
import comfy.model_management
|
||||
device = comfy.model_management.get_torch_device()
|
||||
|
||||
import comfy.ops
|
||||
ops = comfy.ops.manual_cast
|
||||
|
||||
try:
|
||||
import xformers
|
||||
import xformers.ops
|
||||
@@ -77,13 +80,13 @@ class ZeroSFT(nn.Module):
|
||||
nhidden = 128
|
||||
|
||||
self.mlp_shared = nn.Sequential(
|
||||
nn.Conv2d(label_nc, nhidden, kernel_size=ks, padding=pw),
|
||||
ops.Conv2d(label_nc, nhidden, kernel_size=ks, padding=pw),
|
||||
nn.SiLU()
|
||||
)
|
||||
self.zero_mul = zero_module(nn.Conv2d(nhidden, norm_nc + concat_channels, kernel_size=ks, padding=pw))
|
||||
self.zero_add = zero_module(nn.Conv2d(nhidden, norm_nc + concat_channels, kernel_size=ks, padding=pw))
|
||||
# self.zero_mul = nn.Conv2d(nhidden, norm_nc + concat_channels, kernel_size=ks, padding=pw)
|
||||
# self.zero_add = nn.Conv2d(nhidden, norm_nc + concat_channels, kernel_size=ks, padding=pw)
|
||||
self.zero_mul = zero_module(ops.Conv2d(nhidden, norm_nc + concat_channels, kernel_size=ks, padding=pw))
|
||||
self.zero_add = zero_module(ops.Conv2d(nhidden, norm_nc + concat_channels, kernel_size=ks, padding=pw))
|
||||
# self.zero_mul = ops.Conv2d(nhidden, norm_nc + concat_channels, kernel_size=ks, padding=pw)
|
||||
# self.zero_add = ops.Conv2d(nhidden, norm_nc + concat_channels, kernel_size=ks, padding=pw)
|
||||
|
||||
self.zero_conv = zero_module(conv_nd(2, label_nc, norm_nc, 1, 1, 0))
|
||||
self.pre_concat = bool(concat_channels != 0)
|
||||
@@ -296,7 +299,7 @@ class GLVControl(nn.Module):
|
||||
self.label_emb = nn.Embedding(num_classes, time_embed_dim)
|
||||
elif self.num_classes == "continuous":
|
||||
print("setting up linear c_adm embedding layer")
|
||||
self.label_emb = nn.Linear(1, time_embed_dim)
|
||||
self.label_emb = ops.Linear(1, time_embed_dim)
|
||||
elif self.num_classes == "timestep":
|
||||
self.label_emb = checkpoint_wrapper_fn(
|
||||
nn.Sequential(
|
||||
|
||||
@@ -14,9 +14,8 @@ from ..modules.distributions.distributions import DiagonalGaussianDistribution
|
||||
from ..modules.ema import LitEma
|
||||
from ..util import default, get_obj_from_str, instantiate_from_config
|
||||
|
||||
class Conv2d(torch.nn.Conv2d):
|
||||
def reset_parameters(self):
|
||||
return None
|
||||
import comfy.ops
|
||||
ops = comfy.ops.manual_cast
|
||||
|
||||
class AbstractAutoencoder(pl.LightningModule):
|
||||
"""
|
||||
@@ -297,8 +296,8 @@ class AutoencoderKL(AutoencodingEngine):
|
||||
assert ddconfig["double_z"]
|
||||
self.encoder = Encoder(**ddconfig)
|
||||
self.decoder = Decoder(**ddconfig)
|
||||
self.quant_conv = Conv2d(2 * ddconfig["z_channels"], 2 * embed_dim, 1)
|
||||
self.post_quant_conv = Conv2d(embed_dim, ddconfig["z_channels"], 1)
|
||||
self.quant_conv = ops.Conv2d(2 * ddconfig["z_channels"], 2 * embed_dim, 1)
|
||||
self.post_quant_conv = ops.Conv2d(embed_dim, ddconfig["z_channels"], 1)
|
||||
self.embed_dim = embed_dim
|
||||
|
||||
if ckpt_path is not None:
|
||||
|
||||
+29
-34
@@ -10,13 +10,8 @@ from einops import rearrange, repeat
|
||||
from packaging import version
|
||||
from torch import nn
|
||||
|
||||
class Conv2d(torch.nn.Conv2d):
|
||||
def reset_parameters(self):
|
||||
return None
|
||||
|
||||
class Linear(torch.nn.Linear):
|
||||
def reset_parameters(self):
|
||||
return None
|
||||
import comfy.ops
|
||||
ops = comfy.ops.manual_cast
|
||||
|
||||
if version.parse(torch.__version__) >= version.parse("2.0.0"):
|
||||
SDP_IS_AVAILABLE = True
|
||||
@@ -92,7 +87,7 @@ def init_(tensor):
|
||||
class GEGLU(nn.Module):
|
||||
def __init__(self, dim_in, dim_out):
|
||||
super().__init__()
|
||||
self.proj = Linear(dim_in, dim_out * 2)
|
||||
self.proj = ops.Linear(dim_in, dim_out * 2)
|
||||
|
||||
def forward(self, x):
|
||||
x, gate = self.proj(x).chunk(2, dim=-1)
|
||||
@@ -105,13 +100,13 @@ class FeedForward(nn.Module):
|
||||
inner_dim = int(dim * mult)
|
||||
dim_out = default(dim_out, dim)
|
||||
project_in = (
|
||||
nn.Sequential(Linear(dim, inner_dim), nn.GELU())
|
||||
nn.Sequential(ops.Linear(dim, inner_dim), nn.GELU())
|
||||
if not glu
|
||||
else GEGLU(dim, inner_dim)
|
||||
)
|
||||
|
||||
self.net = nn.Sequential(
|
||||
project_in, nn.Dropout(dropout), Linear(inner_dim, dim_out)
|
||||
project_in, nn.Dropout(dropout), ops.Linear(inner_dim, dim_out)
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
@@ -128,7 +123,7 @@ def zero_module(module):
|
||||
|
||||
|
||||
def Normalize(in_channels):
|
||||
return torch.nn.GroupNorm(
|
||||
return ops.GroupNorm(
|
||||
num_groups=32, num_channels=in_channels, eps=1e-6, affine=True
|
||||
)
|
||||
|
||||
@@ -138,8 +133,8 @@ class LinearAttention(nn.Module):
|
||||
super().__init__()
|
||||
self.heads = heads
|
||||
hidden_dim = dim_head * heads
|
||||
self.to_qkv = Conv2d(dim, hidden_dim * 3, 1, bias=False)
|
||||
self.to_out = Conv2d(hidden_dim, dim, 1)
|
||||
self.to_qkv = ops.Conv2d(dim, hidden_dim * 3, 1, bias=False)
|
||||
self.to_out = ops.Conv2d(hidden_dim, dim, 1)
|
||||
|
||||
def forward(self, x):
|
||||
b, c, h, w = x.shape
|
||||
@@ -162,16 +157,16 @@ class SpatialSelfAttention(nn.Module):
|
||||
self.in_channels = in_channels
|
||||
|
||||
self.norm = Normalize(in_channels)
|
||||
self.q = Conv2d(
|
||||
self.q = ops.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
self.k = Conv2d(
|
||||
self.k = ops.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
self.v = Conv2d(
|
||||
self.v = ops.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
self.proj_out = Conv2d(
|
||||
self.proj_out = ops.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
|
||||
@@ -218,12 +213,12 @@ class CrossAttention(nn.Module):
|
||||
self.scale = dim_head**-0.5
|
||||
self.heads = heads
|
||||
|
||||
self.to_q = Linear(query_dim, inner_dim, bias=False)
|
||||
self.to_k = Linear(context_dim, inner_dim, bias=False)
|
||||
self.to_v = Linear(context_dim, inner_dim, bias=False)
|
||||
self.to_q = ops.Linear(query_dim, inner_dim, bias=False)
|
||||
self.to_k = ops.Linear(context_dim, inner_dim, bias=False)
|
||||
self.to_v = ops.Linear(context_dim, inner_dim, bias=False)
|
||||
|
||||
self.to_out = nn.Sequential(
|
||||
Linear(inner_dim, query_dim), nn.Dropout(dropout)
|
||||
ops.Linear(inner_dim, query_dim), nn.Dropout(dropout)
|
||||
)
|
||||
self.backend = backend
|
||||
|
||||
@@ -309,12 +304,12 @@ class MemoryEfficientCrossAttention(nn.Module):
|
||||
self.heads = heads
|
||||
self.dim_head = dim_head
|
||||
|
||||
self.to_q = Linear(query_dim, inner_dim, bias=False)
|
||||
self.to_k = Linear(context_dim, inner_dim, bias=False)
|
||||
self.to_v = Linear(context_dim, inner_dim, bias=False)
|
||||
self.to_q = ops.Linear(query_dim, inner_dim, bias=False)
|
||||
self.to_k = ops.Linear(context_dim, inner_dim, bias=False)
|
||||
self.to_v = ops.Linear(context_dim, inner_dim, bias=False)
|
||||
|
||||
self.to_out = nn.Sequential(
|
||||
Linear(inner_dim, query_dim), nn.Dropout(dropout)
|
||||
ops.Linear(inner_dim, query_dim), nn.Dropout(dropout)
|
||||
)
|
||||
self.attention_op: Optional[Any] = None
|
||||
|
||||
@@ -442,9 +437,9 @@ class BasicTransformerBlock(nn.Module):
|
||||
dropout=dropout,
|
||||
backend=sdp_backend,
|
||||
) # is self-attn if context is none
|
||||
self.norm1 = nn.LayerNorm(dim)
|
||||
self.norm2 = nn.LayerNorm(dim)
|
||||
self.norm3 = nn.LayerNorm(dim)
|
||||
self.norm1 = ops.LayerNorm(dim)
|
||||
self.norm2 = ops.LayerNorm(dim)
|
||||
self.norm3 = ops.LayerNorm(dim)
|
||||
self.checkpoint = checkpoint
|
||||
#if self.checkpoint:
|
||||
#print(f"{self.__class__.__name__} is using checkpointing")
|
||||
@@ -523,8 +518,8 @@ class BasicTransformerSingleLayerBlock(nn.Module):
|
||||
context_dim=context_dim,
|
||||
)
|
||||
self.ff = FeedForward(dim, dropout=dropout, glu=gated_ff)
|
||||
self.norm1 = nn.LayerNorm(dim)
|
||||
self.norm2 = nn.LayerNorm(dim)
|
||||
self.norm1 = ops.LayerNorm(dim)
|
||||
self.norm2 = ops.LayerNorm(dim)
|
||||
self.checkpoint = checkpoint
|
||||
|
||||
def forward(self, x, context=None):
|
||||
@@ -588,11 +583,11 @@ class SpatialTransformer(nn.Module):
|
||||
inner_dim = n_heads * d_head
|
||||
self.norm = Normalize(in_channels)
|
||||
if not use_linear:
|
||||
self.proj_in = Conv2d(
|
||||
self.proj_in = ops.Conv2d(
|
||||
in_channels, inner_dim, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
else:
|
||||
self.proj_in = Linear(in_channels, inner_dim)
|
||||
self.proj_in = ops.Linear(in_channels, inner_dim)
|
||||
|
||||
self.transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
@@ -612,11 +607,11 @@ class SpatialTransformer(nn.Module):
|
||||
)
|
||||
if not use_linear:
|
||||
self.proj_out = zero_module(
|
||||
Conv2d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0)
|
||||
ops.Conv2d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0)
|
||||
)
|
||||
else:
|
||||
# self.proj_out = zero_module(Linear(in_channels, inner_dim))
|
||||
self.proj_out = zero_module(Linear(inner_dim, in_channels))
|
||||
self.proj_out = zero_module(ops.Linear(inner_dim, in_channels))
|
||||
self.use_linear = use_linear
|
||||
|
||||
def forward(self, x, context=None):
|
||||
|
||||
@@ -8,6 +8,8 @@ from torchvision import models
|
||||
|
||||
from ..util import get_ckpt_path
|
||||
|
||||
import comfy.ops
|
||||
ops = comfy.ops.manual_cast
|
||||
|
||||
class LPIPS(nn.Module):
|
||||
# Learned perceptual metric
|
||||
@@ -91,7 +93,7 @@ class NetLinLayer(nn.Module):
|
||||
else []
|
||||
)
|
||||
layers += [
|
||||
nn.Conv2d(chn_in, chn_out, 1, stride=1, padding=0, bias=False),
|
||||
ops.Conv2d(chn_in, chn_out, 1, stride=1, padding=0, bias=False),
|
||||
]
|
||||
self.model = nn.Sequential(*layers)
|
||||
|
||||
|
||||
@@ -3,7 +3,8 @@ import functools
|
||||
import torch.nn as nn
|
||||
|
||||
from ..util import ActNorm
|
||||
|
||||
import comfy.ops
|
||||
ops = comfy.ops.manual_cast
|
||||
|
||||
def weights_init(m):
|
||||
classname = m.__class__.__name__
|
||||
@@ -42,7 +43,7 @@ class NLayerDiscriminator(nn.Module):
|
||||
kw = 4
|
||||
padw = 1
|
||||
sequence = [
|
||||
nn.Conv2d(input_nc, ndf, kernel_size=kw, stride=2, padding=padw),
|
||||
ops.Conv2d(input_nc, ndf, kernel_size=kw, stride=2, padding=padw),
|
||||
nn.LeakyReLU(0.2, True),
|
||||
]
|
||||
nf_mult = 1
|
||||
@@ -51,7 +52,7 @@ class NLayerDiscriminator(nn.Module):
|
||||
nf_mult_prev = nf_mult
|
||||
nf_mult = min(2**n, 8)
|
||||
sequence += [
|
||||
nn.Conv2d(
|
||||
ops.Conv2d(
|
||||
ndf * nf_mult_prev,
|
||||
ndf * nf_mult,
|
||||
kernel_size=kw,
|
||||
@@ -66,7 +67,7 @@ class NLayerDiscriminator(nn.Module):
|
||||
nf_mult_prev = nf_mult
|
||||
nf_mult = min(2**n_layers, 8)
|
||||
sequence += [
|
||||
nn.Conv2d(
|
||||
ops.Conv2d(
|
||||
ndf * nf_mult_prev,
|
||||
ndf * nf_mult,
|
||||
kernel_size=kw,
|
||||
@@ -79,7 +80,7 @@ class NLayerDiscriminator(nn.Module):
|
||||
]
|
||||
|
||||
sequence += [
|
||||
nn.Conv2d(ndf * nf_mult, 1, kernel_size=kw, stride=1, padding=padw)
|
||||
ops.Conv2d(ndf * nf_mult, 1, kernel_size=kw, stride=1, padding=padw)
|
||||
] # output 1 channel prediction map
|
||||
self.main = nn.Sequential(*sequence)
|
||||
|
||||
|
||||
@@ -19,13 +19,8 @@ except:
|
||||
|
||||
from ...modules.attention import LinearAttention, MemoryEfficientCrossAttention
|
||||
|
||||
class Conv2d(torch.nn.Conv2d):
|
||||
def reset_parameters(self):
|
||||
return None
|
||||
|
||||
class Linear(torch.nn.Linear):
|
||||
def reset_parameters(self):
|
||||
return None
|
||||
import comfy.ops
|
||||
ops = comfy.ops.manual_cast
|
||||
|
||||
def get_timestep_embedding(timesteps, embedding_dim):
|
||||
"""
|
||||
@@ -54,7 +49,7 @@ def nonlinearity(x):
|
||||
|
||||
|
||||
def Normalize(in_channels, num_groups=32):
|
||||
return torch.nn.GroupNorm(
|
||||
return ops.GroupNorm(
|
||||
num_groups=num_groups, num_channels=in_channels, eps=1e-6, affine=True
|
||||
)
|
||||
|
||||
@@ -64,7 +59,7 @@ class Upsample(nn.Module):
|
||||
super().__init__()
|
||||
self.with_conv = with_conv
|
||||
if self.with_conv:
|
||||
self.conv = Conv2d(
|
||||
self.conv = ops.Conv2d(
|
||||
in_channels, in_channels, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
|
||||
@@ -91,7 +86,7 @@ class Downsample(nn.Module):
|
||||
self.with_conv = with_conv
|
||||
if self.with_conv:
|
||||
# no asymmetric padding in torch conv, must do it ourselves
|
||||
self.conv = Conv2d(
|
||||
self.conv = ops.Conv2d(
|
||||
in_channels, in_channels, kernel_size=3, stride=2, padding=0
|
||||
)
|
||||
|
||||
@@ -122,23 +117,23 @@ class ResnetBlock(nn.Module):
|
||||
self.use_conv_shortcut = conv_shortcut
|
||||
|
||||
self.norm1 = Normalize(in_channels)
|
||||
self.conv1 = Conv2d(
|
||||
self.conv1 = ops.Conv2d(
|
||||
in_channels, out_channels, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
if temb_channels > 0:
|
||||
self.temb_proj = Linear(temb_channels, out_channels)
|
||||
self.temb_proj = ops.Linear(temb_channels, out_channels)
|
||||
self.norm2 = Normalize(out_channels)
|
||||
self.dropout = torch.nn.Dropout(dropout)
|
||||
self.conv2 = Conv2d(
|
||||
self.conv2 = ops.Conv2d(
|
||||
out_channels, out_channels, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
if self.in_channels != self.out_channels:
|
||||
if self.use_conv_shortcut:
|
||||
self.conv_shortcut = Conv2d(
|
||||
self.conv_shortcut = ops.Conv2d(
|
||||
in_channels, out_channels, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
else:
|
||||
self.nin_shortcut = Conv2d(
|
||||
self.nin_shortcut = ops.Conv2d(
|
||||
in_channels, out_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
|
||||
@@ -178,16 +173,16 @@ class AttnBlock(nn.Module):
|
||||
self.in_channels = in_channels
|
||||
|
||||
self.norm = Normalize(in_channels)
|
||||
self.q = Conv2d(
|
||||
self.q = ops.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
self.k = Conv2d(
|
||||
self.k = ops.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
self.v = Conv2d(
|
||||
self.v = ops.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
self.proj_out = Conv2d(
|
||||
self.proj_out = ops.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
|
||||
@@ -228,16 +223,16 @@ class MemoryEfficientAttnBlock(nn.Module):
|
||||
self.in_channels = in_channels
|
||||
|
||||
self.norm = Normalize(in_channels)
|
||||
self.q = Conv2d(
|
||||
self.q = ops.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
self.k = Conv2d(
|
||||
self.k = ops.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
self.v = Conv2d(
|
||||
self.v = ops.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
self.proj_out = Conv2d(
|
||||
self.proj_out = ops.Conv2d(
|
||||
in_channels, in_channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
self.attention_op: Optional[Any] = None
|
||||
@@ -354,13 +349,13 @@ class Model(nn.Module):
|
||||
self.temb = nn.Module()
|
||||
self.temb.dense = nn.ModuleList(
|
||||
[
|
||||
Linear(self.ch, self.temb_ch),
|
||||
Linear(self.temb_ch, self.temb_ch),
|
||||
ops.Linear(self.ch, self.temb_ch),
|
||||
ops.Linear(self.temb_ch, self.temb_ch),
|
||||
]
|
||||
)
|
||||
|
||||
# downsampling
|
||||
self.conv_in = Conv2d(
|
||||
self.conv_in = ops.Conv2d(
|
||||
in_channels, self.ch, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
|
||||
@@ -439,7 +434,7 @@ class Model(nn.Module):
|
||||
|
||||
# end
|
||||
self.norm_out = Normalize(block_in)
|
||||
self.conv_out = Conv2d(
|
||||
self.conv_out = ops.Conv2d(
|
||||
block_in, out_ch, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
|
||||
@@ -526,7 +521,7 @@ class Encoder(nn.Module):
|
||||
self.in_channels = in_channels
|
||||
|
||||
# downsampling
|
||||
self.conv_in = Conv2d(
|
||||
self.conv_in = ops.Conv2d(
|
||||
in_channels, self.ch, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
|
||||
@@ -577,7 +572,7 @@ class Encoder(nn.Module):
|
||||
|
||||
# end
|
||||
self.norm_out = Normalize(block_in)
|
||||
self.conv_out = Conv2d(
|
||||
self.conv_out = ops.Conv2d(
|
||||
block_in,
|
||||
2 * z_channels if double_z else z_channels,
|
||||
kernel_size=3,
|
||||
@@ -660,7 +655,7 @@ class Decoder(nn.Module):
|
||||
make_resblock_cls = self._make_resblock()
|
||||
make_conv_cls = self._make_conv()
|
||||
# z to block_in
|
||||
self.conv_in = Conv2d(
|
||||
self.conv_in = ops.Conv2d(
|
||||
z_channels, block_in, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
|
||||
@@ -719,7 +714,7 @@ class Decoder(nn.Module):
|
||||
return ResnetBlock
|
||||
|
||||
def _make_conv(self) -> Callable:
|
||||
return Conv2d
|
||||
return ops.Conv2d
|
||||
|
||||
def get_last_layer(self, **kwargs):
|
||||
return self.conv_out.weight
|
||||
|
||||
@@ -19,6 +19,9 @@ import comfy.model_management
|
||||
device = comfy.model_management.get_torch_device()
|
||||
from contextlib import nullcontext
|
||||
|
||||
import comfy.ops
|
||||
ops = comfy.ops.manual_cast
|
||||
|
||||
def make_beta_schedule(
|
||||
schedule,
|
||||
n_timestep,
|
||||
@@ -278,25 +281,17 @@ class GroupNorm32(nn.GroupNorm):
|
||||
def forward(self, x):
|
||||
# return super().forward(x.float()).type(x.dtype)
|
||||
return super().forward(x)
|
||||
|
||||
class Conv2d(torch.nn.Conv2d):
|
||||
def reset_parameters(self):
|
||||
return None
|
||||
|
||||
class Linear(torch.nn.Linear):
|
||||
def reset_parameters(self):
|
||||
return None
|
||||
|
||||
def conv_nd(dims, *args, **kwargs):
|
||||
"""
|
||||
Create a 1D, 2D, or 3D convolution module.
|
||||
"""
|
||||
if dims == 1:
|
||||
return nn.Conv1d(*args, **kwargs)
|
||||
return ops.Conv1d(*args, **kwargs)
|
||||
elif dims == 2:
|
||||
return Conv2d(*args, **kwargs)
|
||||
return ops.Conv2d(*args, **kwargs)
|
||||
elif dims == 3:
|
||||
return nn.Conv3d(*args, **kwargs)
|
||||
return ops.Conv3d(*args, **kwargs)
|
||||
raise ValueError(f"unsupported dimensions: {dims}")
|
||||
|
||||
|
||||
@@ -304,7 +299,7 @@ def linear(*args, **kwargs):
|
||||
"""
|
||||
Create a linear module.
|
||||
"""
|
||||
return Linear(*args, **kwargs)
|
||||
return ops.Linear(*args, **kwargs)
|
||||
|
||||
|
||||
def avg_pool_nd(dims, *args, **kwargs):
|
||||
|
||||
Reference in New Issue
Block a user