Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
111 changes: 48 additions & 63 deletions taufactor/electrode.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,36 +40,28 @@ def __init__(self, img, conductive_label=1, reactive_label=0, \
self.cpu_img = None

def return_mask(self, img):
return (img == self.cond_label).to(dtype=self.precision)
return img == self.cond_label

def init_field(self, mask):
x = np.arange(self.Nx)+0.5
c_init = self.electrode_bc + (self.left_bc-self.electrode_bc)*np.cosh(1-x/self.Nx)/np.cosh(1)
vec = torch.tensor(c_init, dtype=self.precision, device=self.device)
for i in range(2):
vec = torch.unsqueeze(vec, -1)
vec = torch.unsqueeze(vec, 0)
return self._pad(mask * vec, [self.left_bc * 2, 0])
return self._init_field_from_profile(mask, vec, (self.left_bc * 2, 0))

def init_conductive_neighbours(self, img, mask):
padded = self._pad(mask, [2, 0])
cond_nn = torch.empty_like(mask)
self._neighbour_sum_from_padded(padded, cond_nn)
del padded
cond_nn.masked_fill_(mask == 0, torch.inf)
cond_nn = torch.empty(mask.shape, dtype=self.precision, device=self.device)
self._neighbour_sum(mask, cond_nn, (2, 0))
cond_nn.masked_fill_(~mask, torch.inf)
return cond_nn

def init_reactive_neighbours(self, img):
reac = (img == self.reac_label).to(dtype=self.precision)
padded = self._pad(reac)
del reac
# A count of 0..6, so uint8 is ample and 4x smaller than float32
reac_nn = torch.empty(
(self.batch_size, self.Nx, self.Ny, self.Nz),
dtype=self.precision,
dtype=torch.uint8,
device=self.device,
)
self._neighbour_sum_from_padded(padded, reac_nn)
del padded
self._neighbour_sum(img == self.reac_label, reac_nn)
reac_nn.masked_fill_(img != self.cond_label, 0)
return reac_nn

Expand All @@ -81,11 +73,7 @@ def compute_metrics(self):
relative_error = np.max(np.abs(c_x-self.c_x), axis=1)
self.c_x = c_x

fluxes = -self.field[:, 1:-1, 1:-1, 1:-1] + self.field[:, :-2, 1:-1, 1:-1]
fluxes[:, 0, :, :] = (self.left_bc-self.field[:, 1, 1:-1, 1:-1])/0.5
fluxes[self.field[:, 1:-1, 1:-1, 1:-1] == 0] = 0
fluxes[self.field[:, :-2, 1:-1, 1:-1] == 0] = 0
fluxes = torch.mean(fluxes, (2, 3)).cpu().numpy()
fluxes = self._slice_flux()
fluxes_1d = np.concatenate((2*(self.left_bc-c_x[:,:1]), (-c_x[:,1:]+c_x[:,:-1])), axis=1)
fluxes_1d[:,1:][self.vol_x[:,1:]==0] = 0
fluxes_1d[:,1:][self.vol_x[:,:-1]==0] = 0
Expand All @@ -110,6 +98,31 @@ def compute_metrics(self):
self.Z_ideal = compute_impedance_batched(R_ideal, C_ideal, freq)
tau = self.Z_sim[:,0].real/self.Z_ideal[:,0].real
return tau, relative_error

def _slice_flux(self):
"""Mean x-flux per slice, shape (bs, Nx).

Only the (bs, Nx) reduction is needed, so it is accumulated in x chunks
rather than building a full-volume flux field.
"""
out = torch.empty((self.batch_size, self.Nx), dtype=self.field.dtype,
device=self.device)
chunk = self._chunk_size
buf = torch.empty((self.batch_size, chunk, self.Ny, self.Nz),
dtype=self.field.dtype, device=self.device)
for x0 in range(0, self.Nx, chunk):
x1 = min(self.Nx, x0 + chunk)
flux = buf[:, :x1 - x0]
centre = self.field[:, x0 + 1:x1 + 1, 1:-1, 1:-1]
upwind = self.field[:, x0:x1, 1:-1, 1:-1]
torch.sub(upwind, centre, out=flux)
if x0 == 0:
flux[:, 0] = (self.left_bc - self.field[:, 1, 1:-1, 1:-1]) / 0.5
flux.masked_fill_(centre == 0, 0)
flux.masked_fill_(upwind == 0, 0)
out[:, x0:x1] = torch.mean(flux, (2, 3))
return out.cpu().numpy()


def plot_stats(self, relative_error):
clear_output(wait=True)
Expand Down Expand Up @@ -141,24 +154,18 @@ class PeriodicElectrodeSolver(ElectrodeSolver):
Solver with periodic boundary conditions in y and z direction.
"""
def init_conductive_neighbours(self, img, mask):
padded = self._pad(mask, [2, 0])[:, :, 1:-1, 1:-1]
cond_nn = torch.empty_like(mask)
self._periodic_yz_neighbour_sum_from_padded(padded, cond_nn)
del padded
cond_nn.masked_fill_(mask == 0, torch.inf)
cond_nn = torch.empty(mask.shape, dtype=self.precision, device=self.device)
self._periodic_yz_neighbour_sum(mask, cond_nn, (2, 0))
cond_nn.masked_fill_(~mask, torch.inf)
return cond_nn

def init_reactive_neighbours(self, img):
reac = (img == self.reac_label).to(dtype=self.precision)
padded = self._pad(reac)[:, :, 1:-1, 1:-1]
del reac
reac_nn = torch.empty(
(self.batch_size, self.Nx, self.Ny, self.Nz),
dtype=self.precision,
dtype=torch.uint8,
device=self.device,
)
self._periodic_yz_neighbour_sum_from_padded(padded, reac_nn)
del padded
self._periodic_yz_neighbour_sum(img == self.reac_label, reac_nn)
reac_nn.masked_fill_(img != self.cond_label, 0)
return reac_nn

Expand Down Expand Up @@ -217,7 +224,7 @@ def __init__(self, img, conductive_label=1, reactive_label=0, \
self.c_x = 0

def return_mask(self, img):
return (img == self.cond_label).to(dtype=self.precision)
return img == self.cond_label

def init_field(self, img):
return None
Expand All @@ -240,11 +247,7 @@ def init_field_internal(self, mask):
vec = torch.tensor(phi, dtype=torch.complex64, device=self.device)
else:
vec = torch.tensor(phi, dtype=torch.complex128, device=self.device)
for i in range(2):
vec = torch.unsqueeze(vec, -1)
vec = torch.unsqueeze(vec, 0)
mask_f = mask.to(vec.dtype) if mask.dtype == torch.bool else mask
return self._pad(mask_f * vec, [2 * self.left_bc, 0]).to(self.device)
return self._init_field_from_profile(mask, vec, (2 * self.left_bc, 0))

def count_neighbours(self, img, mask):
"""
Expand All @@ -259,18 +262,12 @@ def count_neighbours(self, img, mask):
:rtype: cp.array
"""
# Conducting nearest neighbours
padded = self._pad(mask, [2, 0])
cond_nn = torch.empty_like(mask)
self._neighbour_sum_from_padded(padded, cond_nn)
del padded
self._neighbour_sum(mask, cond_nn, (2, 0))

# Capacitive nearest neighbours
reac = (img == self.reac_label).to(dtype=self.precision)
padded = self._pad(reac)
del reac
reac_nn = torch.empty_like(mask)
self._neighbour_sum_from_padded(padded, reac_nn)
del padded
self._neighbour_sum(img == self.reac_label, reac_nn)

# Masking conducting voxels
cond_nn[mask == 0] = 0.0
Expand Down Expand Up @@ -419,23 +416,11 @@ class PeriodicImpedanceSolver(ImpedanceSolver):
Solver with periodic boundary conditions in y and z direction.
"""
def count_neighbours(self, img, mask):
img2 = self._pad(mask, [2, 0])[:, :, 1:-1, 1:-1]
cond_nn = torch.zeros_like(img2)
# iterate through shifts in the spatial dimensions
for dim in range(1, 4):
for dr in [1, -1]:
cond_nn += torch.roll(img2, dr, dim)
cond_nn = cond_nn[:, 1:-1]

img2 = torch.zeros_like(img)
img2[img==self.reac_label] = 1
img2 = self._pad(img2)[:, :, 1:-1, 1:-1]
reac_nn = torch.zeros_like(img2)
# iterate through shifts in the spatial dimensions
for dim in range(1, 4):
for dr in [1, -1]:
reac_nn += torch.roll(img2, dr, dim)
reac_nn = reac_nn[:, 1:-1]
cond_nn = torch.empty_like(mask)
self._periodic_yz_neighbour_sum(mask, cond_nn, (2, 0))

reac_nn = torch.empty_like(mask)
self._periodic_yz_neighbour_sum(img == self.reac_label, reac_nn)

# Masking conducting voxels
cond_nn[mask == 0] = 0.0
Expand Down
Loading
Loading