mirror of
https://github.com/wassname/PointCNN.git
synced 2026-09-06 16:30:36 +08:00
Fixing inconsistencies. Model up to ~94% on MNIST. Subtle further tuning likely required.
This commit is contained in:
+150
-41
@@ -1,4 +1,7 @@
|
||||
import os
|
||||
|
||||
import math
|
||||
import random
|
||||
import data_utils
|
||||
import time
|
||||
|
||||
@@ -7,10 +10,17 @@ import torch
|
||||
from torch import nn
|
||||
from torch.autograd import Variable
|
||||
from torch.utils.data import Dataset, DataLoader
|
||||
import matplotlib.pyplot as plt
|
||||
from mpl_toolkits.mplot3d import Axes3D
|
||||
|
||||
from pointcnn.core import rPointCNN
|
||||
from pointcnn.util import knn_indices_func_gpu
|
||||
from pointcnn.layers import Dense
|
||||
from visualize import *
|
||||
|
||||
random.seed(0)
|
||||
|
||||
CURRENT_DIR = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
x = 2
|
||||
|
||||
@@ -44,29 +54,57 @@ class Classifier(nn.Module):
|
||||
def __init__(self):
|
||||
super(Classifier, self).__init__()
|
||||
|
||||
self.pcnn = nn.Sequential(
|
||||
paPointCNN( 1, 32, 8, 1, 256),
|
||||
paPointCNN( 32, 64, 8, 2, 256),
|
||||
paPointCNN( 64, 96, 8, 4, 256),
|
||||
self.pcnn1 = paPointCNN( 1, 32, 8, 1, -1)
|
||||
self.pcnn2 = nn.Sequential(
|
||||
paPointCNN( 32, 64, 8, 2, -1),
|
||||
paPointCNN( 64, 96, 8, 4, -1),
|
||||
paPointCNN( 96, 128, 12, 4, 120),
|
||||
paPointCNN(128, 160, 12, 6, 120),
|
||||
paPointCNN(128, 160, 12, 6, 120)
|
||||
)
|
||||
|
||||
self.fcn = nn.Sequential(
|
||||
Dense(160, 128),
|
||||
Dense(128, 64, drop_rate = 0.5),
|
||||
Dense(64, 10, activation = None)
|
||||
Dense( 64, 10, with_bn = False, activation = None)
|
||||
)
|
||||
|
||||
self.log_softmax = nn.LogSoftmax()
|
||||
|
||||
def forward(self, x):
|
||||
x = self.pcnn(x)[1] # grab features
|
||||
x = self.pcnn1(x)
|
||||
if False:
|
||||
print("Making graph...")
|
||||
k = make_dot(x[1])
|
||||
|
||||
print("Viewing...")
|
||||
k.view()
|
||||
print("DONE")
|
||||
|
||||
assert False
|
||||
x = self.pcnn2(x)[1] # grab features
|
||||
|
||||
logits = self.fcn(x)
|
||||
# logits = torch.mean(logits, dim = 1)
|
||||
return logits
|
||||
# log_probs = self.log_softmax(logits)
|
||||
# return log_probs
|
||||
logits_mean = torch.mean(logits, dim = 1)
|
||||
return logits_mean
|
||||
|
||||
"""
|
||||
def get_indices(batch_size, sample_num, point_num, random_sample = True):
|
||||
indices = []
|
||||
for i in range(batch_size):
|
||||
|
||||
if random_sample:
|
||||
# point_num >= sample_num generally
|
||||
choices = np.random.choice(point_num, sample_num, replace = (point_num < sample_num))
|
||||
else:
|
||||
# This modulo generally not used.
|
||||
choices = np.arange(sample_num) % point_num
|
||||
|
||||
choices = np.expand_dims(choices, axis = 0)
|
||||
b_idx_mat = np.full_like(choices, i)
|
||||
|
||||
# Each set of choices is paired with its batch index.
|
||||
choices_2d = np.concatenate((b_idx_max, choices), axis = 0)
|
||||
indices.append(choices_2d)
|
||||
return np.stack(indices, axis = 1) # (2, batch_size,
|
||||
"""
|
||||
|
||||
model = Classifier().cuda()
|
||||
|
||||
@@ -95,58 +133,129 @@ batch_num = batch_num_per_epoch * num_epochs
|
||||
training_set = mnist_dataset(data_train, label_train)
|
||||
training_loader = DataLoader(training_set, batch_size = batch_size)
|
||||
|
||||
testing_batch_size = 256
|
||||
testing_set = mnist_dataset(data_val, label_val)
|
||||
testing_loader = DataLoader(testing_set, batch_size = 1)
|
||||
testing_loader = DataLoader(testing_set, batch_size = testing_batch_size)
|
||||
|
||||
lr = 0.01
|
||||
decay_steps = 8000
|
||||
decay_rate = 0.6
|
||||
lr_min = 0.00001
|
||||
|
||||
optimizer = torch.optim.SGD(model.parameters(), lr = 0.01, momentum = 0.9)
|
||||
loss_fn = nn.NLLLoss()
|
||||
loss_fn = nn.CrossEntropyLoss()
|
||||
|
||||
for _ in range(num_epochs):
|
||||
global_step = 1
|
||||
|
||||
n = 0
|
||||
model_save_dir = os.path.join(CURRENT_DIR, "models", "mnist2")
|
||||
os.makedirs(model_save_dir, exist_ok = True)
|
||||
|
||||
losses = []
|
||||
accuracies = []
|
||||
|
||||
if False:
|
||||
latest_model = sorted(os.listdir(model_save_dir))[-1]
|
||||
model.load_state_dict(torch.load(os.path.join(model_save_dir, latest_model)))
|
||||
|
||||
for e in range(1, num_epochs + 1):
|
||||
|
||||
print("EPOCH %i of %i" % (e, num_epochs))
|
||||
|
||||
m_loc = os.path.join(model_save_dir, "save_e%.4d" % e)
|
||||
torch.save(model.state_dict(), m_loc)
|
||||
np.savez_compressed(os.path.join(CURRENT_DIR, "losses"), losses)
|
||||
np.savez_compressed(os.path.join(CURRENT_DIR, "accuracies"), accuracies)
|
||||
|
||||
# Decaying learning rate
|
||||
if e > 1:
|
||||
lr *= decay_rate ** (global_step // decay_steps)
|
||||
if lr > lr_min:
|
||||
print("NEW LEARNING RATE:", lr)
|
||||
optimizer = torch.optim.SGD(model.parameters(), lr = lr, momentum = 0.9)
|
||||
|
||||
for data, label in training_loader:
|
||||
|
||||
n += 1
|
||||
model = model.train()
|
||||
|
||||
data = Variable(data).cuda()
|
||||
label = Variable(label.long()).cuda()
|
||||
label = label.long()
|
||||
P = data[:,:,:3]
|
||||
F = data[:,:,3:]
|
||||
|
||||
offset = int(random.gauss(0, sample_num // 8))
|
||||
offset = max(offset, -sample_num // 4)
|
||||
offset = min(offset, sample_num // 4)
|
||||
sample_num_train = sample_num + offset
|
||||
|
||||
# indices = get_indices(batch_size, sample_num_train, point_num)
|
||||
indices = np.random.choice(P.size()[1], sample_num_train, replace = False).tolist()
|
||||
P_sampled = P[:,indices,:]
|
||||
F_sampled = F[:,indices,:]
|
||||
P_sampled = Variable(P_sampled).cuda()
|
||||
F_sampled = Variable(F_sampled).cuda()
|
||||
|
||||
if False:
|
||||
P_draw = P_sampled.data.cpu().numpy()
|
||||
# fig = plt.figure()
|
||||
# ax = fig.gca(projection = '3d')
|
||||
# ax.scatter(P_draw[25,:,0], P_draw[25,:,1], P_draw[25,:,2], c = 'k')
|
||||
# plt.show()
|
||||
|
||||
plt.style.use('grayscale')
|
||||
plt.axis([-3, 3, -3, 3])
|
||||
plt.scatter(P_draw[25,:,0], -P_draw[25,:,1], c = 1 - F_sampled[25,:,0], marker = ',', s = 25)
|
||||
print("LABEL:", label[25])
|
||||
plt.show()
|
||||
|
||||
optimizer.zero_grad()
|
||||
|
||||
t0 = time.time()
|
||||
out = model((P, F))
|
||||
out = model((P_sampled, F_sampled))
|
||||
|
||||
print(out)
|
||||
|
||||
loss = loss_fn(out, label)
|
||||
loss = loss_fn(out, Variable(label.long()).cuda())
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
|
||||
print("loss:", loss.data[0])
|
||||
if global_step % 25 == 0:
|
||||
loss_v = loss.data[0]
|
||||
print("Loss:", loss_v)
|
||||
else:
|
||||
loss_v = 0
|
||||
|
||||
if global_step % 250 == 0:
|
||||
|
||||
if n % 25 == 0:
|
||||
# Testing accuracy
|
||||
num_testing = 0
|
||||
total = 0
|
||||
correct = 0
|
||||
for data, label in testing_loader:
|
||||
if num_testing > 100:
|
||||
accuracy_sum = 0
|
||||
testing_size = 4 # times testing_batch_size = 256
|
||||
for t, (data, label) in enumerate(testing_loader):
|
||||
if t >= testing_size:
|
||||
break
|
||||
else:
|
||||
num_testing += 1
|
||||
data = Variable(data).cuda()
|
||||
label = Variable(label.long()).cuda()
|
||||
|
||||
model = model.eval()
|
||||
|
||||
P = data[:,:,:3]
|
||||
F = data[:,:,3:]
|
||||
out = model((P, F))
|
||||
|
||||
offset = int(random.gauss(0, sample_num // 8))
|
||||
offset = max(offset, -sample_num // 4)
|
||||
offset = min(offset, sample_num // 4)
|
||||
sample_num_train = sample_num + offset
|
||||
|
||||
# indices = get_indices(batch_size, sample_num_train, point_num)
|
||||
indices = np.random.choice(P.size()[1], sample_num_train, replace = False).tolist()
|
||||
P_sampled = P[:,indices,:]
|
||||
F_sampled = F[:,indices,:]
|
||||
P_sampled = Variable(P_sampled).cuda()
|
||||
F_sampled = Variable(F_sampled).cuda()
|
||||
|
||||
out = model((P_sampled, F_sampled))
|
||||
probs = nn.Softmax()(out)
|
||||
# print(probs)
|
||||
_, pred = probs.max(1)
|
||||
total += 1
|
||||
if pred.cpu().data[0] == label.cpu().data[0]:
|
||||
correct += 1
|
||||
accuracy = correct / total
|
||||
print(pred)
|
||||
accuracy_sum += torch.mean((pred.data.cpu() == label.long()).float())
|
||||
accuracy = accuracy_sum / testing_size
|
||||
print("accuracy:", accuracy)
|
||||
|
||||
losses.append(loss_v)
|
||||
accuracies.append(accuracy)
|
||||
|
||||
global_step += 1
|
||||
|
||||
+88
-44
@@ -8,11 +8,11 @@ import matplotlib.pyplot as plt
|
||||
|
||||
try:
|
||||
from .util import knn_indices_func, knn_indices_func_gpu
|
||||
from .layers import MLP, LayerNorm, Conv, SepConv, endchannels
|
||||
from .layers import MLP, LayerNorm, Conv, SepConv, Dense, endchannels
|
||||
# from .context import timed
|
||||
except SystemError:
|
||||
from util import knn_indices_func, knn_indices_func_gpu
|
||||
from layers import MLP, LayerNorm, Conv, SepConv, endchannels
|
||||
from layers import MLP, LayerNorm, Conv, SepConv, Dense, endchannels
|
||||
# from context import timed
|
||||
|
||||
class XConv(nn.Module):
|
||||
@@ -20,8 +20,7 @@ class XConv(nn.Module):
|
||||
Vectorized pointwise convolution.
|
||||
"""
|
||||
|
||||
def __init__(self, C_in, C_out, D, N_neighbors, N_rep, C_lifted = None,
|
||||
mlp_width = 2):
|
||||
def __init__(self, C_in, C_out, D, N_neighbors, N_rep, C_lifted, depth_multiplier):
|
||||
"""
|
||||
:param C_in: Input dimension of the points' features.
|
||||
:param C_out: Output dimension of the representative point features.
|
||||
@@ -32,9 +31,6 @@ class XConv(nn.Module):
|
||||
"""
|
||||
super(XConv, self).__init__()
|
||||
|
||||
if C_lifted == None:
|
||||
C_lifted = C_in # Not optimal?
|
||||
|
||||
if __debug__:
|
||||
# Only needed for assertions.
|
||||
self.C_in = C_in
|
||||
@@ -48,9 +44,22 @@ class XConv(nn.Module):
|
||||
# self.pts_layernorm = LayerNorm(2, momentum = 0.9)
|
||||
|
||||
# Main dense linear layers
|
||||
self.mlp_lift = MLP([D] + [self.C_lifted] * (mlp_width - 1), batch_norm = False)
|
||||
self.dense1 = Dense(D, C_lifted)
|
||||
self.dense2 = Dense(C_lifted, C_lifted)
|
||||
|
||||
# Layers to generate X
|
||||
self.x_trans = nn.Sequential(
|
||||
endchannels(Conv(
|
||||
in_channels = D,
|
||||
out_channels = N_neighbors**2,
|
||||
kernel_size = (1, N_neighbors),
|
||||
with_bn = False
|
||||
)),
|
||||
Dense(N_neighbors**2, N_neighbors**2, with_bn = False),
|
||||
Dense(N_neighbors**2, N_neighbors**2, with_bn = False, activation = None)
|
||||
)
|
||||
|
||||
"""
|
||||
self.mid_conv = endchannels(Conv(D, N_neighbors**2, (1, N_neighbors))).cuda()
|
||||
self.mid_dwconv1 = endchannels(SepConv(
|
||||
in_channels = N_neighbors,
|
||||
@@ -64,14 +73,15 @@ class XConv(nn.Module):
|
||||
kernel_size = (1, N_neighbors),
|
||||
depth_multiplier = N_neighbors
|
||||
)).cuda()
|
||||
"""
|
||||
|
||||
print(depth_multiplier)
|
||||
# Final
|
||||
self.mlp = MLP([N_neighbors] * mlp_width, batch_norm = False)
|
||||
self.end_conv = endchannels(SepConv(
|
||||
in_channels = C_lifted + C_in,
|
||||
out_channels = C_out,
|
||||
kernel_size = (1, N_neighbors),
|
||||
depth_multiplier = 4
|
||||
depth_multiplier = depth_multiplier
|
||||
)).cuda()
|
||||
|
||||
# @timed.timed
|
||||
@@ -88,11 +98,16 @@ class XConv(nn.Module):
|
||||
:return: Features aggregated into point p.
|
||||
"""
|
||||
p, P, F = x
|
||||
assert(p.size()[0] == P.size()[0] == F.size()[0]) # Check N is equal.
|
||||
assert(p.size()[1] == P.size()[1] == F.size()[1]) # Check N_rep is equal.
|
||||
assert(P.size()[2] == F.size()[2] == self.N_neighbors) # Check N_neighbors is equal.
|
||||
assert(p.size()[2] == P.size()[3] == self.D) # Check D is equal.
|
||||
assert(F.size()[3] == self.C_in) # Check C_in is equal.
|
||||
if F is not None:
|
||||
assert(p.size()[0] == P.size()[0] == F.size()[0]) # Check N is equal.
|
||||
assert(p.size()[1] == P.size()[1] == F.size()[1]) # Check N_rep is equal.
|
||||
assert(P.size()[2] == F.size()[2] == self.N_neighbors) # Check N_neighbors is equal.
|
||||
assert(F.size()[3] == self.C_in) # Check C_in is equal.
|
||||
else:
|
||||
assert(p.size()[0] == P.size()[0]) # Check N is equal.
|
||||
assert(p.size()[1] == P.size()[1]) # Check N_rep is equal.
|
||||
assert(P.size()[2] == self.N_neighbors) # Check N_neighbors is equal.
|
||||
assert(p.size()[2] == P.size()[3] == self.D) # Check D is equal.
|
||||
|
||||
N = len(P)
|
||||
N_rep = p.size()[1]
|
||||
@@ -103,19 +118,28 @@ class XConv(nn.Module):
|
||||
# P_local = self.pts_layernorm(P - p_center)
|
||||
|
||||
# Individually lift each point into C_lifted dim space.
|
||||
F_lifted = self.mlp_lift(P_local)
|
||||
F_lifted0 = self.dense1(P_local)
|
||||
F_lifted = self.dense2(F_lifted0)
|
||||
|
||||
# Cat F_lifted and F, to size (N, N_rep, N_neighbors, C_lifted + C_in).
|
||||
F_cat = torch.cat((F_lifted, F), -1)
|
||||
# Cat F_lifted and F,None to size (N, N_rep, N_neighbors, C_lifted + C_in).
|
||||
if F is None:
|
||||
F_cat = F_lifted
|
||||
else:
|
||||
F_cat = torch.cat((F_lifted, F), -1)
|
||||
|
||||
# Learn the (N, K, K) X-transformation matrix.
|
||||
X_shape = (N, N_rep, self.N_neighbors, self.N_neighbors)
|
||||
X = self.x_trans(P_local)
|
||||
X = X.view(*X_shape)
|
||||
|
||||
"""
|
||||
X = self.mid_conv(P_local)
|
||||
X = X.contiguous().view(*X_shape)
|
||||
X = self.mid_dwconv1(X)
|
||||
X = X.contiguous().view(*X_shape)
|
||||
X = self.mid_dwconv2(X)
|
||||
X = X.contiguous().view(*X_shape)
|
||||
"""
|
||||
|
||||
# Weight and permute F_cat with the learned X.
|
||||
F_X = torch.matmul(X, F_cat)
|
||||
@@ -127,8 +151,7 @@ class PointCNN(nn.Module):
|
||||
TODO: Insert documentation
|
||||
"""
|
||||
|
||||
def __init__(self, C_in, C_out, D, N_neighbors, dilution, N_rep,
|
||||
r_indices_func, C_lifted = None, mlp_width = 2):
|
||||
def __init__(self, C_in, C_out, D, N_neighbors, dilution, N_rep, r_indices_func):
|
||||
"""
|
||||
:param C_in: Input dimension of the points' features.
|
||||
:param C_out: Output dimension of the representative point features.
|
||||
@@ -156,11 +179,12 @@ class PointCNN(nn.Module):
|
||||
"""
|
||||
super(PointCNN, self).__init__()
|
||||
|
||||
if C_lifted == None:
|
||||
C_lifted = C_in # Not optimal?
|
||||
C_lifted = C_out // 2 if C_in == 0 else C_out // 4
|
||||
depth_multiplier = min(int(np.ceil(C_out / C_in)), 4)
|
||||
|
||||
self.r_indices_func = r_indices_func
|
||||
self.x_conv = XConv(C_in, C_out, D, N_neighbors, N_rep, C_lifted, mlp_width)
|
||||
self.dense = Dense(C_in, C_out // 2) if C_in != 0 else None
|
||||
self.x_conv = XConv(C_out // 2 if C_in != 0 else C_in, C_out, D, N_neighbors, N_rep, C_lifted, depth_multiplier)
|
||||
self.dilution = dilution
|
||||
|
||||
def select_region(self, P, P_idx):
|
||||
@@ -193,6 +217,8 @@ class PointCNN(nn.Module):
|
||||
:return:
|
||||
"""
|
||||
ps, P, F = x
|
||||
t0 = time.time()
|
||||
F = self.dense(F) if F is not None else F
|
||||
P_idx = self.r_indices_func(ps, P, self.x_conv.N_neighbors, self.dilution) # This step takes ~97% of the time.
|
||||
P_regional = self.select_region(P, P_idx) # Prime target for optimization: KNN on GPU.
|
||||
if False:
|
||||
@@ -205,9 +231,12 @@ class PointCNN(nn.Module):
|
||||
plt.scatter(test_point[0], test_point[1], s = 100, c = 'green')
|
||||
plt.scatter(neighborhood[:,0], neighborhood[:,1], s = 100, c = 'red')
|
||||
plt.show()
|
||||
F_regional = self.select_region(F, P_idx)
|
||||
# ps, P, F_P -> ps_F
|
||||
t1 = time.time()
|
||||
F_regional = self.select_region(F, P_idx) if F is not None else F
|
||||
test_time = time.time() - t1
|
||||
F_p = self.x_conv((ps, P_regional, F_regional))
|
||||
total_time = time.time() - t0
|
||||
# print("frac of time:", test_time / total_time)
|
||||
return F_p
|
||||
|
||||
class rPointCNN(nn.Module):
|
||||
@@ -220,7 +249,7 @@ class rPointCNN(nn.Module):
|
||||
|
||||
def forward(self, x):
|
||||
P, F = x
|
||||
if self.N_rep < P.size()[1]:
|
||||
if 0 < self.N_rep < P.size()[1]:
|
||||
idx = np.random.choice(P.size()[1], self.N_rep, replace = False).tolist()
|
||||
ps = P[:,idx,:]
|
||||
else:
|
||||
@@ -229,30 +258,45 @@ class rPointCNN(nn.Module):
|
||||
ps_F = self.pointcnn((ps, P, F))
|
||||
return ps, ps_F
|
||||
|
||||
def plot(P, F):
|
||||
num_F = F.size()[2]
|
||||
pts = P[0].data.cpu().numpy()
|
||||
plt.scatter(pts[:,0], pts[:,1], s = num_F, c = "k")
|
||||
plt.savefig("./%i.png" % num_F)
|
||||
plt.cla()
|
||||
|
||||
if __name__ == "__main__":
|
||||
np.random.seed(0)
|
||||
|
||||
N = 1
|
||||
num_points = 500
|
||||
N_rep = 20
|
||||
num_points = 1000
|
||||
D = 2
|
||||
C_in = 16
|
||||
C_out = 32
|
||||
N_neighbors = 30
|
||||
C_in = 4
|
||||
N_neighbors = 10
|
||||
dilution = 1
|
||||
|
||||
model = PointCNN(C_in, C_out, D, N_neighbors, dilution, N_rep, knn_indices_func_gpu).cuda()
|
||||
layer1 = rPointCNN(C_in, 8, D, N_neighbors, dilution, 1000, knn_indices_func).cuda()
|
||||
layer2 = rPointCNN( 8, 16, D, N_neighbors, dilution, 500, knn_indices_func).cuda()
|
||||
layer3 = rPointCNN( 16, 32, D, N_neighbors, dilution, 250, knn_indices_func).cuda()
|
||||
layer4 = rPointCNN( 32, 64, D, N_neighbors, dilution, 125, knn_indices_func).cuda()
|
||||
layer5 = rPointCNN( 64, 128, D, N_neighbors, dilution, 50, knn_indices_func).cuda()
|
||||
|
||||
test_P = np.random.rand(N,num_points,D).astype(np.float32)
|
||||
test_F = np.random.rand(N,num_points,C_in).astype(np.float32)
|
||||
idx = np.random.choice(test_P.shape[1], N_rep, replace = False)
|
||||
test_ps = test_P[:,idx,:]
|
||||
P = np.random.rand(N,num_points,D).astype(np.float32)
|
||||
F = np.random.rand(N,num_points,C_in).astype(np.float32)
|
||||
P = Variable(torch.from_numpy(P)).cuda()
|
||||
F = Variable(torch.from_numpy(F)).cuda()
|
||||
|
||||
test_P = Variable(torch.from_numpy(test_P)).cuda()
|
||||
test_F = Variable(torch.from_numpy(test_F)).cuda()
|
||||
test_ps = Variable(torch.from_numpy(test_ps)).cuda()
|
||||
|
||||
print(test_F.size())
|
||||
for _ in range(1):
|
||||
out = model((test_ps, test_P, test_F))
|
||||
print(out.size())
|
||||
if True:
|
||||
P, F = layer1((P, F))
|
||||
else:
|
||||
plot(P, F)
|
||||
P, F = layer1((P, F))
|
||||
plot(P, F)
|
||||
P, F = layer2((P, F))
|
||||
plot(P, F)
|
||||
P, F = layer3((P, F))
|
||||
plot(P, F)
|
||||
P, F = layer4((P, F))
|
||||
plot(P, F)
|
||||
P, F = layer5((P, F))
|
||||
plot(P, F)
|
||||
|
||||
+2
-2
@@ -50,7 +50,7 @@ class Conv(nn.Module):
|
||||
|
||||
self.conv = nn.Conv2d(in_channels, out_channels, kernel_size, bias = not with_bn)
|
||||
self.activation = activation
|
||||
self.bn = nn.BatchNorm2d(out_channels) if with_bn else None
|
||||
self.bn = nn.BatchNorm2d(out_channels, momentum = 0.9) if with_bn else None
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv(x)
|
||||
@@ -75,7 +75,7 @@ class SepConv(nn.Module):
|
||||
)
|
||||
|
||||
self.activation = activation
|
||||
self.bn = nn.BatchNorm2d(out_channels) if with_bn else None
|
||||
self.bn = nn.BatchNorm2d(out_channels, momentum = 0.9) if with_bn else None
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv(x)
|
||||
|
||||
Reference in New Issue
Block a user