fuse-element 0.1.dev0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
fuse/cells.py ADDED
@@ -0,0 +1,850 @@
1
+ import matplotlib as mpl
2
+ mpl.use('Agg')
3
+ import matplotlib.pyplot as plt
4
+ import numpy as np
5
+ import itertools
6
+ import networkx as nx
7
+ import fuse.groups as fe_groups
8
+ import copy
9
+ import sympy as sp
10
+ from matplotlib.patches import FancyArrowPatch
11
+ from mpl_toolkits.mplot3d import proj3d
12
+ from sympy.combinatorics.named_groups import SymmetricGroup
13
+ from fuse.utils import sympy_to_numpy, fold_reduce
14
+ from FIAT.reference_element import Simplex, UFCQuadrilateral
15
+ from ufl.cell import Cell
16
+
17
+
18
+ class Arrow3D(FancyArrowPatch):
19
+ def __init__(self, xs, ys, zs, *args, **kwargs):
20
+ super().__init__((0, 0), (0, 0), *args, **kwargs)
21
+ self._verts3d = xs, ys, zs
22
+
23
+ def do_3d_projection(self, renderer=None):
24
+ xs3d, ys3d, zs3d = self._verts3d
25
+ xs, ys, zs = proj3d.proj_transform(xs3d, ys3d, zs3d, self.axes.M)
26
+ self.set_positions((xs[0], ys[0]), (xs[1], ys[1]))
27
+
28
+ return np.min(zs)
29
+
30
+
31
+ def topo_pos(G):
32
+ """
33
+ Helper function for hasse diagram visualisation
34
+ Offsets the nodes and displays in topological order
35
+ """
36
+ pos_dict = {}
37
+ for i, node_list in enumerate(nx.topological_generations(G)):
38
+ x_offset = len(node_list) / 2
39
+ y_offset = 0
40
+ for j, name in enumerate(node_list):
41
+ pos_dict[name] = (j - x_offset, i - j * y_offset)
42
+
43
+ return pos_dict
44
+
45
+
46
+ def normalise(v):
47
+ norm = np.linalg.norm(v)
48
+ return v / norm
49
+
50
+
51
+ def make_arrow(ax, mid, edge, direction=1):
52
+ delta = 0.0001 if direction >= 0 else -0.0001
53
+ x, y = edge(mid)
54
+ dir_x, dir_y = edge(mid + delta)
55
+ ax.arrow(x, y, dir_x-x, dir_y-y, head_width=0.05, head_length=0.1)
56
+
57
+
58
+ def make_arrow_3d(ax, mid, edge, direction=1):
59
+ delta = 0.0001 if direction >= 0 else -0.0001
60
+ x, y, z = edge(mid)
61
+ dir_x, dir_y, dir_z = edge(mid + delta)
62
+ a = Arrow3D([x, dir_x], [y, dir_y], [z, dir_z], mutation_scale=10, arrowstyle="-|>", color="black")
63
+ ax.add_artist(a)
64
+
65
+
66
+ def construct_attach_2d(a, b, c, d):
67
+ """
68
+ Compute polynomial attachment in x based on two points (a,b) and (c,d)
69
+
70
+ :param: a,b,c,d: two points (a,b) and (c,d)
71
+ """
72
+ x = sp.Symbol("x")
73
+ return [((c-a)/2)*(x+1) + a, ((d-b)/2)*(x+1) + b]
74
+
75
+
76
+ def construct_attach_3d(res):
77
+ """
78
+ Convert matrix of coefficients into a vector of polynomials in x and y
79
+
80
+ :param: res: matrix of coefficients
81
+ """
82
+ x = sp.Symbol("x")
83
+ y = sp.Symbol("y")
84
+ xy = sp.Matrix([1, x, y])
85
+ return (xy.T * res)
86
+
87
+
88
+ def compute_scaled_verts(d, n):
89
+ """
90
+ Construct default cell vertices
91
+
92
+ :param: d: dimension of cell
93
+ :param: n: number of vertices
94
+ """
95
+ if d == 2:
96
+ source = np.array([0, 1])
97
+ rot_coords = [source for i in range(0, n)]
98
+
99
+ rot_mat = np.array([[np.cos((2*np.pi)/n), -np.sin((2*np.pi)/n)], [np.sin((2*np.pi)/n), np.cos((2*np.pi)/n)]])
100
+ for i in range(1, n):
101
+ rot_coords[i] = np.matmul(rot_mat, rot_coords[i-1])
102
+ xdiff, ydiff = (rot_coords[0][0] - rot_coords[1][0],
103
+ rot_coords[0][1] - rot_coords[1][1])
104
+ scale = 2 / np.sqrt(xdiff**2 + ydiff**2)
105
+ scaled_coords = np.array([[scale*x, scale*y] for (x, y) in rot_coords])
106
+ return scaled_coords
107
+ elif d == 3:
108
+ if n == 4:
109
+ A = [-1, 1, -1]
110
+ B = [1, -1, -1]
111
+ C = [1, 1, 1]
112
+ D = [-1, -1, 1]
113
+ coords = [A, B, C, D]
114
+ face1 = np.array([A, D, C])
115
+ face2 = np.array([A, B, D])
116
+ face3 = np.array([A, C, B])
117
+ face4 = np.array([B, D, C])
118
+ faces = [face1, face2, face3, face4]
119
+ elif n == 8:
120
+ coords = []
121
+ faces = [[] for i in range(6)]
122
+ for i in [-1, 1]:
123
+ for j in [-1, 1]:
124
+ for k in [-1, 1]:
125
+ coords.append([i, j, k])
126
+
127
+ for j in [-1, 1]:
128
+ for k in [-1, 1]:
129
+ faces[0].append([1, j, k])
130
+ faces[1].append([-1, j, k])
131
+ faces[2].append([j, 1, k])
132
+ faces[3].append([j, -1, k])
133
+ faces[4].append([j, k, 1])
134
+ faces[5].append([j, k, -1])
135
+
136
+ else:
137
+ raise ValueError("Polyhedron with {} vertices not supported".format(n))
138
+
139
+ xdiff, ydiff, zdiff = (coords[0][0] - coords[1][0],
140
+ coords[0][1] - coords[1][1],
141
+ coords[0][2] - coords[1][2])
142
+ scale = 2 / np.sqrt(xdiff**2 + ydiff**2 + zdiff**2)
143
+ scaled_coords = np.array([[scale*x, scale*y, scale*z] for (x, y, z) in coords])
144
+ scaled_faces = np.array([[[scale*x, scale*y, scale*z] for (x, y, z) in face] for face in faces])
145
+
146
+ return scaled_coords, scaled_faces
147
+ else:
148
+ raise ValueError("Dimension {} not supported".format(d))
149
+
150
+
151
+ def polygon(n):
152
+ """
153
+ Constructs the 2D default cell with n sides/vertices
154
+
155
+ :param: n: number of vertices
156
+ """
157
+ vertices = []
158
+ for i in range(n):
159
+ vertices.append(Point(0))
160
+ edges = []
161
+ for i in range(n):
162
+ edges.append(
163
+ Point(1, [vertices[(i+1) % n], vertices[(i+2) % n]], vertex_num=2))
164
+
165
+ return Point(2, edges, vertex_num=n)
166
+
167
+
168
+ def make_tetrahedron():
169
+ vertices = []
170
+ for i in range(4):
171
+ vertices.append(Point(0))
172
+ edges = []
173
+ edges.append(
174
+ Point(1, vertex_num=2, edges=[vertices[0], vertices[1]]))
175
+ edges.append(
176
+ Point(1, vertex_num=2, edges=[vertices[1], vertices[2]]))
177
+ edges.append(
178
+ Point(1, vertex_num=2, edges=[vertices[2], vertices[0]]))
179
+ edges.append(
180
+ Point(1, vertex_num=2, edges=[vertices[3], vertices[0]]))
181
+ edges.append(
182
+ Point(1, vertex_num=2, edges=[vertices[1], vertices[3]]))
183
+ edges.append(
184
+ Point(1, vertex_num=2, edges=[vertices[2], vertices[3]]))
185
+
186
+ face1 = Point(2, vertex_num=3, edges=[edges[5], edges[3], edges[2]], edge_orientations={2: [1, 0]})
187
+ face2 = Point(2, vertex_num=3, edges=[edges[3], edges[0], edges[4]])
188
+ face3 = Point(2, vertex_num=3, edges=[edges[2], edges[0], edges[1]])
189
+ face4 = Point(2, vertex_num=3, edges=[edges[1], edges[4], edges[5]], edge_orientations={0: [1, 0], 2: [1, 0]})
190
+
191
+ return Point(3, vertex_num=4, edges=[face3, face1, face4, face2])
192
+
193
+
194
+ class Point():
195
+ """
196
+ Cell complex representation of a finite element cell
197
+
198
+ :param: d: dimension of the cell
199
+ :param: edges: list of subcells (either as edge or point objects)
200
+ :param: vertex_num: Optional argument, number of vertices
201
+ :param: oriented: adds orientation to the cell
202
+ :param: group: Symmetry group of the cell
203
+ :param: edge_orientations: dictionary of the orientations of the subcells
204
+
205
+ """
206
+
207
+ id_iter = itertools.count()
208
+
209
+ def __init__(self, d, edges=[], vertex_num=None, oriented=False, group=None, edge_orientations={}, cell_id=None):
210
+ if not cell_id:
211
+ cell_id = next(self.id_iter)
212
+ self.id = cell_id
213
+ self.dimension = d
214
+ if d == 0:
215
+ assert (edges == [])
216
+ if vertex_num:
217
+ edges = self.compute_attachments(vertex_num, edges, edge_orientations)
218
+
219
+ self.oriented = oriented
220
+ self.G = nx.DiGraph()
221
+ self.G.add_node(self.id, point_class=self)
222
+ for edge in edges:
223
+ assert edge.lower_dim() < self.dimension
224
+ self.G.add_edge(self.id, edge.point.id, edge_class=edge)
225
+ self.G = nx.compose_all([self.G]
226
+ + [edge.point.graph() for edge in edges])
227
+ self.connections = edges
228
+
229
+ self.group = group
230
+ if not group:
231
+ self.group = self.compute_cell_group()
232
+
233
+ self.group = self.group.add_cell(self)
234
+
235
+ def compute_attachments(self, n, points, orientations={}):
236
+ """
237
+ Compute the attachment function between two nodes
238
+
239
+ :param: n: number of vertices
240
+ :param: points: List of Point objects
241
+ :param: orientations: (Optional) Orientation associated with the attachment
242
+ """
243
+ if self.dimension == 1:
244
+ edges = [Edge(points[0], sp.sympify((-1,))),
245
+ Edge(points[1], sp.sympify((1,)))]
246
+ if self.dimension == 2:
247
+ coords = compute_scaled_verts(2, n)
248
+ edges = []
249
+
250
+ for i in range(n):
251
+ a, b = coords[i]
252
+ c, d = coords[(i + 1) % n]
253
+
254
+ if i in orientations.keys():
255
+ edges.append(Edge(points[i], construct_attach_2d(a, b, c, d), o=points[i].group.get_member(orientations[i])))
256
+ else:
257
+ edges.append(Edge(points[i], construct_attach_2d(a, b, c, d)))
258
+ if self.dimension == 3:
259
+ coords, faces = compute_scaled_verts(3, n)
260
+ coords_2d = np.c_[np.ones(len(faces[0])), compute_scaled_verts(2, len(faces[0]))]
261
+ res = []
262
+ edges = []
263
+
264
+ for i in range(len(faces)):
265
+ res = np.linalg.solve(coords_2d, faces[i])
266
+
267
+ res_fn = construct_attach_3d(res)
268
+ # breakpoint()
269
+ assert np.allclose(np.array(res_fn.subs({"x": coords_2d[0][1], "y": coords_2d[0][2]})).astype(np.float64), faces[i][0])
270
+ assert np.allclose(np.array(res_fn.subs({"x": coords_2d[1][1], "y": coords_2d[1][2]})).astype(np.float64), faces[i][1])
271
+ assert np.allclose(np.array(res_fn.subs({"x": coords_2d[2][1], "y": coords_2d[2][2]})).astype(np.float64), faces[i][2])
272
+ if i in orientations.keys():
273
+ edges.append(Edge(points[i], construct_attach_3d(res), o=points[i].group.get_member(orientations[i])))
274
+ else:
275
+ edges.append(Edge(points[i], construct_attach_3d(res)))
276
+
277
+ # breakpoint()
278
+ return edges
279
+
280
+ def compute_cell_group(self):
281
+ """
282
+ Systematically work out the symmetry group of the constructed cell
283
+ """
284
+ verts = self.ordered_vertices()
285
+ v_coords = [self.get_node(v, return_coords=True) for v in verts]
286
+ n = len(verts)
287
+ max_group = SymmetricGroup(n)
288
+ edges = [edge.ordered_vertices() for edge in self.edges()]
289
+ accepted_perms = max_group.elements.copy()
290
+ if n > 2:
291
+ for element in max_group.elements:
292
+ reordered = element(verts)
293
+ for edge in edges:
294
+ diff = np.subtract(v_coords[reordered.index(edge[0])], v_coords[reordered.index(edge[1])]).squeeze()
295
+ edge_len = np.sqrt(np.dot(diff, diff))
296
+ if not np.allclose(edge_len, 2):
297
+ accepted_perms.remove(element)
298
+ break
299
+ return fe_groups.PermutationSetRepresentation(list(accepted_perms))
300
+
301
+ def get_spatial_dimension(self):
302
+ return self.dimension
303
+
304
+ def dim(self):
305
+ return self.dimension
306
+
307
+ def get_shape(self):
308
+ num_verts = len(self.vertices())
309
+ if num_verts == 1:
310
+ # Point
311
+ return 0
312
+ elif num_verts == 2:
313
+ # Line
314
+ return 1
315
+ elif num_verts == 3:
316
+ # Triangle
317
+ return 2
318
+ elif num_verts == 4:
319
+ if self.dimension == 2:
320
+ # quadrilateral
321
+ return 11
322
+ elif self.dimension == 3:
323
+ # tetrahedron
324
+ return 3
325
+ elif num_verts == 8:
326
+ # hexahedron
327
+ return 111
328
+ else:
329
+ raise TypeError("Shape undefined for {}".format(str(self)))
330
+
331
+ def get_topology(self):
332
+ structure = [sorted(generation) for generation in nx.topological_generations(self.graph())]
333
+ structure.reverse()
334
+
335
+ min_ids = [min(dimension) for dimension in structure]
336
+ vertices = self.ordered_vertices()
337
+ relabelled_verts = {vertices[i]: i for i in range(len(vertices))}
338
+ self.topology = {}
339
+ self.topology_verts = {}
340
+ for i in range(len(structure)):
341
+ dimension = structure[i]
342
+ self.topology[i] = {}
343
+ self.topology_verts[i] = {}
344
+ for node in dimension:
345
+ neighbours = list(self.G.neighbors(node))
346
+ # self.topology_verts[i][node - min_ids[i]] = tuple([vert - min_ids[0] for vert in self.get_node(node).ordered_vertices()])
347
+ self.topology_verts[i][node - min_ids[i]] = tuple([relabelled_verts[vert] for vert in self.get_node(node).ordered_vertices()])
348
+ if len(neighbours) > 0:
349
+ renumbered_neighbours = tuple([neighbour - min_ids[i-1] for neighbour in neighbours])
350
+ self.topology[i][node - min_ids[i]] = renumbered_neighbours
351
+ else:
352
+ self.topology[i][node - min_ids[i]] = (node - min_ids[i], )
353
+ return self.topology_verts
354
+
355
+ def get_starter_ids(self):
356
+ structure = [sorted(generation) for generation in nx.topological_generations(self.G)]
357
+ structure.reverse()
358
+
359
+ min_ids = [min(dimension) for dimension in structure]
360
+ return min_ids
361
+
362
+ def graph_dim(self):
363
+ if self.oriented:
364
+ dim = self.dimension + 1
365
+ else:
366
+ dim = self.dimension
367
+ return dim
368
+
369
+ def graph(self):
370
+ if self.oriented:
371
+ temp_G = self.G.copy()
372
+ temp_G.remove_node(-1)
373
+ return temp_G
374
+ return self.G
375
+
376
+ def hasse_diagram(self, counter=0, filename=None):
377
+ ax = plt.axes()
378
+ nx.draw_networkx(self.G, pos=topo_pos(self.G),
379
+ with_labels=True, ax=ax)
380
+ edge_dict = {(u, v): self.G.edges[u, v]["edge_class"].o for (u, v) in self.G.edges()}
381
+ nx.draw_networkx_edge_labels(self.G, pos=topo_pos(self.G), edge_labels=edge_dict, ax=ax)
382
+ if filename:
383
+ ax.figure.savefig(filename)
384
+ else:
385
+ plt.show()
386
+
387
+ def ordered_vertices(self, get_class=False):
388
+ # define a points vertex order by combining the order of the sub elements
389
+ # vertex list and removing duplicates
390
+ if self.dimension == 0:
391
+ if get_class:
392
+ return [self]
393
+ return [self.id]
394
+ else:
395
+ # convert to dict to remove duplicates while maintaining order
396
+ full_list = [c.ordered_vertices(get_class) for c in self.connections]
397
+ flatten = itertools.chain.from_iterable(full_list)
398
+ verts = list(dict.fromkeys(flatten))
399
+ if self.oriented:
400
+ # make sure this is necessary
401
+ return self.oriented.permute(verts)
402
+ return verts
403
+
404
+ def d_entities_ids(self, d):
405
+ return self.d_entities(d, get_class=False)
406
+
407
+ def d_entities(self, d, get_class=True):
408
+ levels = [sorted(generation)
409
+ for generation in nx.topological_generations(self.G)]
410
+ if get_class:
411
+ res = [self.G.nodes.data("point_class")[i] for i in levels[self.graph_dim() - d]]
412
+ else:
413
+ res = levels[self.graph_dim() - d]
414
+ return res
415
+
416
+ def get_node(self, node, return_coords=False):
417
+ if return_coords:
418
+ top_level_node = self.d_entities_ids(self.graph_dim())[0]
419
+ if self.dimension == 0:
420
+ return [()]
421
+ return self.attachment(top_level_node, node)()
422
+ return self.G.nodes.data("point_class")[node]
423
+
424
+ def dim_of_node(self, node):
425
+ levels = [sorted(generation)
426
+ for generation in nx.topological_generations(self.G)]
427
+ for i in range(len(levels)):
428
+ if node in levels[i]:
429
+ return self.graph_dim() - i
430
+ raise "Error: Node not found in graph"
431
+
432
+ def vertices(self, get_class=True, return_coords=False):
433
+ # TODO maybe refactor with get_node
434
+ verts = self.d_entities(0, get_class)
435
+ if return_coords:
436
+ verts = self.d_entities_ids(0)
437
+ top_level_node = self.d_entities_ids(self.graph_dim())[0]
438
+ if self.dimension == 0:
439
+ return [()]
440
+ return [self.attachment(top_level_node, v)() for v in verts]
441
+ return verts
442
+
443
+ def edges(self, get_class=True):
444
+ return self.d_entities(1, get_class)
445
+
446
+ def permute_entities(self, g, d):
447
+ # TODO something is wrong here for squares it can return [()]
448
+ verts = self.vertices(get_class=False)
449
+ entities = self.d_entities_ids(d)
450
+ reordered = g.permute(verts)
451
+
452
+ if d == 0:
453
+ entity_group = self.d_entities(d)[0].group
454
+ return list(zip(reordered, [entity_group.identity for r in reordered]))
455
+
456
+ entity_dict = {}
457
+ reordered_entity_dict = {}
458
+
459
+ for e in self.d_entities(d):
460
+ entity_dict[e.id] = tuple(e.ordered_vertices())
461
+ reordered_entity_dict[e.id] = tuple([reordered[verts.index(i)] for i in e.ordered_vertices()])
462
+
463
+ reordered_entities = [tuple() for e in range(len(entities))]
464
+ min_id = min(entities)
465
+ entity_group = self.d_entities(d)[0].group
466
+ for ent in entities:
467
+ for ent1 in entities:
468
+ if set(entity_dict[ent]) == set(reordered_entity_dict[ent1]):
469
+ if entity_dict[ent] != reordered_entity_dict[ent1]:
470
+ o = entity_group.transform_between_perms(entity_dict[ent], reordered_entity_dict[ent1])
471
+ reordered_entities[ent1 - min_id] = (ent, o)
472
+ else:
473
+ reordered_entities[ent1 - min_id] = (ent, entity_group.identity)
474
+
475
+ return reordered_entities
476
+
477
+ def basis_vectors(self, return_coords=True, entity=None):
478
+ if not entity:
479
+ entity = self
480
+ entity_levels = [sorted(generation) for generation in nx.topological_generations(entity.G)]
481
+ self_levels = [sorted(generation) for generation in nx.topological_generations(self.G)]
482
+ vertices = entity_levels[entity.graph_dim()]
483
+ if self.dimension == 0:
484
+ # return [[]
485
+ raise ValueError("Dimension 0 entities cannot have Basis Vectors")
486
+ top_level_node = self_levels[0][0]
487
+ v_0 = vertices[0]
488
+ if return_coords:
489
+ v_0_coords = self.attachment(top_level_node, v_0)()
490
+ basis_vecs = []
491
+ for v in vertices[1:]:
492
+ if return_coords:
493
+ v_coords = self.attachment(top_level_node, v)()
494
+ sub = normalise(np.subtract(v_coords, v_0_coords))
495
+ if not hasattr(sub, "__iter__"):
496
+ basis_vecs.append((sub,))
497
+ else:
498
+ basis_vecs.append(tuple(sub))
499
+ else:
500
+ basis_vecs.append((v, v_0))
501
+ return basis_vecs
502
+
503
+ def plot(self, show=True, plain=False, ax=None, filename=None):
504
+ """ for now into 2 dimensional space """
505
+
506
+ top_level_node = self.d_entities(self.graph_dim(), get_class=False)[0]
507
+ xs = np.linspace(-1, 1, 20)
508
+ if ax is None:
509
+ ax = plt.gca()
510
+
511
+ if self.dimension == 1:
512
+ # line plot in 1D case
513
+ nodes = self.d_entities(0, get_class=False)
514
+ points = []
515
+ for node in nodes:
516
+ attach = self.attachment(top_level_node, node)
517
+ points.extend(attach())
518
+ plt.plot(np.array(points), np.zeros_like(points), color="black")
519
+
520
+ for i in range(self.dimension - 1, -1, -1):
521
+ nodes = self.d_entities(i, get_class=False)
522
+ vert_coords = []
523
+ for node in nodes:
524
+ attach = self.attachment(top_level_node, node)
525
+ if i == 0:
526
+ plotted = attach()
527
+ if len(plotted) < 2:
528
+ plotted = (plotted[0], 0)
529
+ vert_coords += [plotted]
530
+ if not plain:
531
+ plt.plot(plotted[0], plotted[1], 'bo')
532
+ plt.annotate(node, (plotted[0], plotted[1]))
533
+ elif i == 1:
534
+ edgevals = np.array([attach(x) for x in xs])
535
+ if len(edgevals[0]) < 2:
536
+ plt.plot(edgevals[:, 0], 0, color="black")
537
+ else:
538
+ plt.plot(edgevals[:, 0], edgevals[:, 1], color="black")
539
+ if not plain:
540
+ make_arrow(ax, 0, attach)
541
+ else:
542
+ raise ValueError("General plotting not implemented")
543
+ # if i == 2:
544
+ # if len(vert_coords) > 2:
545
+ # hull = ConvexHull(vert_coords)
546
+ # plt.fill(vert_coords[hull.vertices, 0], vert_coords[hull.vertices, 1], alpha=0.5)
547
+ if show:
548
+ plt.show()
549
+ if filename:
550
+ ax.figure.savefig(filename)
551
+
552
+ def plot3d(self, show=True, ax=None):
553
+ assert self.dimension == 3
554
+ if ax is None:
555
+ fig = plt.figure()
556
+ ax = fig.add_subplot(projection='3d')
557
+ xs = np.linspace(-1, 1, 20)
558
+
559
+ top_level_node = self.d_entities_ids(self.graph_dim())[0]
560
+ nodes = self.d_entities_ids(0)
561
+ for node in nodes:
562
+ attach = self.attachment(top_level_node, node)
563
+ plotted = attach()
564
+ ax.scatter(plotted[0], plotted[1], plotted[2], color="black")
565
+
566
+ nodes = self.d_entities_ids(1)
567
+ for node in nodes:
568
+ attach = self.attachment(top_level_node, node)
569
+ edgevals = np.array([attach(x) for x in xs])
570
+ ax.plot3D(edgevals[:, 0], edgevals[:, 1], edgevals[:, 2], color="black")
571
+ make_arrow_3d(ax, 0, attach)
572
+ if show:
573
+ plt.show()
574
+
575
+ def attachment(self, source, dst):
576
+ if source == dst:
577
+ # return x
578
+ return lambda *x: x
579
+
580
+ paths = nx.all_simple_edge_paths(self.G, source, dst)
581
+ attachments = [[self.G[s][d]["edge_class"]
582
+ for (s, d) in path] for path in paths]
583
+
584
+ if len(attachments) == 0:
585
+ raise ValueError("No paths from node {} to node {}"
586
+ .format(source, dst))
587
+
588
+ # check all attachments resolve to the same function
589
+ if len(attachments) > 1:
590
+ dst_dim = self.dim_of_node(dst)
591
+ basis = np.eye(dst_dim)
592
+ if dst_dim == 0:
593
+ vals = [fold_reduce(attachment) for attachment in attachments]
594
+ assert all(np.isclose(val, vals[0]).all() for val in vals)
595
+ else:
596
+ for i in range(dst_dim):
597
+ vals = [fold_reduce(attachment, *tuple(basis[i].tolist()))
598
+ for attachment in attachments]
599
+ assert all(np.isclose(val, vals[0]).all() for val in vals)
600
+
601
+ return lambda *x: fold_reduce(attachments[0], *x)
602
+
603
+ def cell_attachment(self, dst):
604
+ if not isinstance(dst, int):
605
+ raise ValueError
606
+ top_level_node = self.d_entities_ids(self.graph_dim())[0]
607
+ return self.attachment(top_level_node, dst)
608
+
609
+ def orient(self, o):
610
+ """ Orientation node is always labelled with -1 """
611
+ oriented_point = copy.deepcopy(self)
612
+ top_level_node = oriented_point.d_entities_ids(
613
+ oriented_point.dimension)[0]
614
+ oriented_point.G.add_node(-1, point_class=None)
615
+ oriented_point.G.add_edge(-1, top_level_node,
616
+ edge_class=Edge(None, o=o))
617
+ oriented_point.oriented = o
618
+ return oriented_point
619
+
620
+ def __repr__(self):
621
+ entity_name = ["v", "e", "f", "c"]
622
+ return entity_name[self.dimension] + str(self.id)
623
+
624
+ def copy(self):
625
+ return copy.deepcopy(self)
626
+
627
+ def to_fiat(self, name=None):
628
+ if len(self.get_topology()[self.dimension][0]) == self.dimension + 1:
629
+ return CellComplexToFiatSimplex(self, name)
630
+ raise NotImplementedError("Non-Simplex elements are not yet supported")
631
+ return CellComplexToFiatCell(self, name)
632
+
633
+ def to_ufl(self, name=None):
634
+ return CellComplexToUFL(self, name)
635
+
636
+ def _to_dict(self):
637
+ # think this is probably missing stuf
638
+ o_dict = {"dim": self.dimension,
639
+ "edges": [c for c in self.connections],
640
+ "oriented": self.oriented,
641
+ "id": self.id}
642
+ return o_dict
643
+
644
+ def dict_id(self):
645
+ return "Cell"
646
+
647
+ def _from_dict(o_dict):
648
+ return Point(o_dict["dim"], o_dict["edges"], oriented=o_dict["oriented"], cell_id=o_dict["id"])
649
+
650
+
651
+ class Edge():
652
+ """
653
+ Representation of the connections in a cell complex.
654
+
655
+ :param: point: the point being connected (lower level)
656
+ :param: attachment: the function describing how the point is attached
657
+ :param: o: orientation function (optional)
658
+ """
659
+
660
+ def __init__(self, point, attachment=None, o=None):
661
+ self.attachment = attachment
662
+ self.point = point
663
+ self.o = o
664
+
665
+ def __call__(self, *x):
666
+ if self.o:
667
+ x = self.o(x)
668
+ if self.attachment:
669
+ syms = ["x", "y", "z"]
670
+ if hasattr(self.attachment, '__iter__'):
671
+ res = []
672
+ for attach_comp in self.attachment:
673
+ if len(attach_comp.atoms(sp.Symbol)) == len(x):
674
+ res.append(sympy_to_numpy(attach_comp, syms, x))
675
+ else:
676
+ res.append(attach_comp.subs({syms[i]: x[i] for i in range(len(x))}))
677
+ return tuple(res)
678
+ return sympy_to_numpy(self.attachment, syms, x)
679
+ return x
680
+
681
+ def ordered_vertices(self, get_class=False):
682
+ verts = self.point.ordered_vertices(get_class)
683
+ if self.o:
684
+ verts = self.o.permute(verts)
685
+ return verts
686
+
687
+ def lower_dim(self):
688
+ return self.point.dim()
689
+
690
+ def __repr__(self):
691
+ return str(self.point)
692
+
693
+ def _to_dict(self):
694
+ o_dict = {"attachment": self.attachment,
695
+ "point": self.point,
696
+ "orientation": self.o}
697
+ return o_dict
698
+
699
+ def dict_id(self):
700
+ return "Edge"
701
+
702
+ def _from_dict(o_dict):
703
+ return Edge(o_dict["point"], o_dict["attachment"], o_dict["orientation"])
704
+
705
+
706
+ class CellComplexToFiatSimplex(Simplex):
707
+ """
708
+ Convert cell complex to fiat
709
+
710
+ :param: cell: a fuse cell complex
711
+
712
+ Currently assumes simplex.
713
+ """
714
+
715
+ def __init__(self, cell, name=None):
716
+ self.fe_cell = cell
717
+ if name is not None:
718
+ name = "IndiaDefCell"
719
+ self.name = name
720
+
721
+ verts = cell.vertices(return_coords=True)
722
+ topology = cell.get_topology()
723
+ shape = cell.get_shape()
724
+ super(CellComplexToFiatSimplex, self).__init__(shape, verts, topology)
725
+
726
+ def cellname(self):
727
+ return self.name
728
+
729
+ def construct_subelement(self, dimension):
730
+ """Constructs the reference element of a cell
731
+ specified by subelement dimension.
732
+
733
+ :arg dimension: subentity dimension (integer)
734
+ """
735
+ return self.fe_cell.d_entities(dimension)[0].to_fiat()
736
+
737
+ def get_facet_element(self):
738
+ dimension = self.get_spatial_dimension()
739
+ return self.construct_subelement(dimension - 1)
740
+
741
+
742
+ class CellComplexToFiatCell(UFCQuadrilateral):
743
+ """
744
+ Convert cell complex to fiat
745
+
746
+ :param: cell: a fuse cell complex
747
+
748
+ Currently assumes simplex.
749
+ """
750
+
751
+ def __init__(self, cell, name=None):
752
+ self.fe_cell = cell
753
+ if name is not None:
754
+ name = "IndiaDefCell"
755
+ self.name = name
756
+
757
+ verts = cell.vertices(return_coords=True)
758
+ topology = cell.get_topology()
759
+ shape = cell.get_shape()
760
+ super(CellComplexToFiatCell, self).__init__(shape, verts, topology)
761
+
762
+ def cellname(self):
763
+ return self.name
764
+
765
+ def construct_subelement(self, dimension):
766
+ """Constructs the reference element of a cell
767
+ specified by subelement dimension.
768
+
769
+ :arg dimension: subentity dimension (integer)
770
+ """
771
+ return self.fe_cell.d_entities(dimension)[0].to_fiat()
772
+
773
+ def get_facet_element(self):
774
+ dimension = self.get_spatial_dimension()
775
+ return self.construct_subelement(dimension - 1)
776
+
777
+ def get_dimension(self):
778
+ return self.get_spatial_dimension()
779
+
780
+
781
+ class CellComplexToUFL(Cell):
782
+ """
783
+ Convert cell complex to UFL
784
+
785
+ :param: cell: a fuse cell complex
786
+
787
+ Currently just maps to a subset of existing UFL cells
788
+ TODO work out generic way around the naming issue
789
+ """
790
+
791
+ def __init__(self, cell, name=None):
792
+ self.cell_complex = cell
793
+
794
+ # TODO work out generic way around the naming issue
795
+ if not name:
796
+ num_verts = len(cell.vertices())
797
+ if num_verts == 1:
798
+ # Point
799
+ name = "vertex"
800
+ elif num_verts == 2:
801
+ # Line
802
+ name = "interval"
803
+ elif num_verts == 3:
804
+ # Triangle
805
+ name = "triangle"
806
+ elif num_verts == 4:
807
+ if cell.dimension == 2:
808
+ # quadrilateral
809
+ name = "quadrilateral"
810
+ elif cell.dimension == 3:
811
+ # tetrahedron
812
+ name = "tetrahedron"
813
+ elif num_verts == 8:
814
+ # hexahedron
815
+ name = "hexahedron"
816
+ else:
817
+ raise TypeError("UFL cell conversion undefined for {}".format(str(cell)))
818
+ super(CellComplexToUFL, self).__init__(name)
819
+
820
+ def to_fiat(self):
821
+ return self.cell_complex.to_fiat(name=self.cellname())
822
+
823
+ def __repr__(self):
824
+ return super(CellComplexToUFL, self).__repr__() + " Complex"
825
+
826
+ def reconstruct(self, **kwargs):
827
+ """Reconstruct this cell, overwriting properties by those in kwargs."""
828
+ cell = self.cell_complex
829
+ for key, value in kwargs.items():
830
+ if key == "cell":
831
+ cell = value
832
+ else:
833
+ raise TypeError(f"reconstruct() got unexpected keyword argument '{key}'")
834
+ return CellComplexToUFL(cell, self._cellname)
835
+
836
+
837
+ def constructCellComplex(name):
838
+ if name == "vertex":
839
+ return Point(0).to_ufl(name)
840
+ elif name == "interval":
841
+ return Point(1, [Point(0), Point(0)], vertex_num=2).to_ufl(name)
842
+ elif name == "triangle":
843
+ return polygon(3).to_ufl(name)
844
+ elif name == "quadrilateral":
845
+ # return Cell(name)
846
+ return polygon(4).to_ufl(name)
847
+ elif name == "tetrahedron":
848
+ return make_tetrahedron().to_ufl(name)
849
+ else:
850
+ raise TypeError("Cell complex construction undefined for {}".format(str(name)))