import timeit import jax import jax.numpy as jnp import matplotlib.pyplot as plt from scimba_jax.linear_approximation.meshes.mesh import Mesh from scimba_jax.linear_approximation.quad.gauss_quad import UnitSquareTensorized from scimba_jax.mapping.mapping import InvertibleFunction, Mapping def plot_mesh(mesh, show_indices_faces=True, show_indices_cells=True, direction=None): """Affiche le maillage 2D ou 3D avec les indices des (d-1)-cells. Si direction est fourni (0, 1 ou 2), affiche uniquement les faces perpendiculaires à cet axe, colorées en bleu. """ dim = mesh.dim assert dim in [2, 3], ( "Seuls les maillages 2D et 3D sont supportés pour l'affichage." ) assert direction is None or 0 <= direction < dim, ( f"direction doit être None ou un entier dans [0, {dim - 1}]." ) n_cells = [int(n) for n in mesh.n_cells] # 1. Préparation des coordonnées du maillage # Création de la grille de référence [0, 1]^d grids = [jnp.linspace(0, 1, n + 1) for n in n_cells] mesh_grids = jnp.meshgrid(*grids, indexing="ij") # Empilement et application du mapping physique nodes_ref = jnp.stack([g.ravel() for g in mesh_grids], axis=-1) nodes_phys = mesh.mapping.local_mapping(nodes_ref) # Reshape pour accès par indices : (n1+1, n2+1, ..., dim) nodes = nodes_phys.reshape(*(n + 1 for n in n_cells), dim) # 2. Configuration de la figure fig = plt.figure(figsize=(10, 8)) if dim == 2: ax = fig.add_subplot(111) colors = ["blue", "red"] # Type 0, Type 1 elif dim == 3: ax = fig.add_subplot(111, projection="3d") colors = ["blue", "red", "green"] # Type 0, Type 1, Type 2 else: print(f"Dimension {dim} non supportée.") return # 3. Boucle sur les types de faces (i-invariant faces) dirs_to_plot = [direction] if direction is not None else range(dim) for i in dirs_to_plot: stride = mesh.faces_stride[i] color = colors[direction] if direction is not None else colors[i] # Dimensions autres que 'i' other_dims = [d for d in range(dim) if d != i] # On itère sur toutes les faces de ce type # La boucle suit l'ordre du reshape de ta fonction get_face_indices for idx_i in range(n_cells[i] + 1): # On génère toutes les combinaisons pour les autres dimensions # Pour simplifier l'affichage, on utilise des boucles imbriquées if dim == 2: for j in range(n_cells[other_dims[0]]): # Points de l'arête p_idx1 = [0, 0] p_idx1[i], p_idx1[other_dims[0]] = idx_i, j p_idx2 = list(p_idx1) p_idx2[other_dims[0]] += 1 p1, p2 = nodes[tuple(p_idx1)], nodes[tuple(p_idx2)] # Dessin ax.plot( [p1[0], p2[0]], [p1[1], p2[1]], color=color, lw=1, alpha=0.7 ) if show_indices_faces: local_idx = idx_i * n_cells[other_dims[0]] + j mid = (p1 + p2) / 2 ax.text( mid[0], mid[1], str(stride + local_idx), color=color, fontsize=8, ha="center", ) elif dim == 3: for j in range(n_cells[other_dims[0]]): for k in range(n_cells[other_dims[1]]): # Coins de la face (quadrilatère) c = [0, 0, 0] c[i], c[other_dims[0]], c[other_dims[1]] = idx_i, j, k p1 = nodes[tuple(c)] c[other_dims[0]] += 1 p2 = nodes[tuple(c)] c[other_dims[1]] += 1 p3 = nodes[tuple(c)] c[other_dims[0]] -= 1 p4 = nodes[tuple(c)] # Dessin contour pts = jnp.array([p1, p2, p3, p4, p1]) ax.plot( pts[:, 0], pts[:, 1], pts[:, 2], color=color, lw=0.5, alpha=0.5, ) if show_indices_faces: # Calcul d'indice cohérent avec le reshape (i, j, k) local_idx = ( idx_i * (n_cells[other_dims[0]] * n_cells[other_dims[1]]) + j * n_cells[other_dims[1]] + k ) mid = (p1 + p2 + p3 + p4) / 4 ax.text( mid[0], mid[1], mid[2], str(stride + local_idx), color=color, fontsize=7, ) # 4. Affichage des indices de cellules if show_indices_cells: if dim == 2: for i in range(n_cells[0]): for j in range(n_cells[1]): p = ( nodes[i, j] + nodes[i + 1, j] + nodes[i + 1, j + 1] + nodes[i, j + 1] ) / 4 cell_idx = i * n_cells[1] + j ax.text( p[0], p[1], str(cell_idx), color="black", fontsize=9, ha="center", va="center", ) elif dim == 3: for i in range(n_cells[0]): for j in range(n_cells[1]): for k in range(n_cells[2]): p = ( nodes[i, j, k] + nodes[i + 1, j, k] + nodes[i + 1, j + 1, k] + nodes[i, j + 1, k] + nodes[i, j, k + 1] + nodes[i + 1, j, k + 1] + nodes[i + 1, j + 1, k + 1] + nodes[i, j + 1, k + 1] ) / 8 cell_idx = i * n_cells[1] * n_cells[2] + j * n_cells[2] + k ax.text( p[0], p[1], p[2], str(cell_idx), color="black", fontsize=8, ha="center", va="center", ) # 5. Affichage des nœuds if dim == 2: ax.scatter(nodes_phys[:, 0], nodes_phys[:, 1], s=5, color="black", zorder=5) ax.set_aspect("equal") ax.set_xlabel("x") ax.set_ylabel("y") else: ax.scatter( nodes_phys[:, 0], nodes_phys[:, 1], nodes_phys[:, 2], s=2, color="black" ) ax.set_xlabel("x") ax.set_ylabel("y") ax.set_zlabel("z") axis_names = ["x", "y", "z"] dir_str = f" - faces axe {axis_names[direction]}" if direction is not None else "" plt.title(f"Maillage {dim}D - {tuple(n_cells)} cellules{dir_str}") plt.show() mapping_id = InvertibleFunction(lambda x: x, lambda y: y) print("\n\n@@@@@@@@@@@@@@@@ test 1d mesh (3,) @@@@@@@@@@@@@@@") m = Mesh( dim=1, n_cells=(3,), ref_quad=UnitSquareTensorized(dim=1, order=3), mapping=Mapping([mapping_id]), ) for i in range(m.n_cells_total): assert i == m._midx_to_fidx(m._fidx_to_midx(i)) assert jnp.all(m.n_i_faces == jnp.array([4])) assert m.n_faces == 4 assert jnp.all(m.faces_stride == jnp.array([0, 4])) def midx_to_fidx(x: jnp.ndarray): return m._midx_to_fidx_s(x) midx_to_fidx_jitted = jax.jit(midx_to_fidx) assert jnp.all( m._midx_to_fidx_s(jnp.array([[-1], [m.n_cells[0]]])) == jnp.array([-1, -2]) ) assert jnp.all( midx_to_fidx_jitted(jnp.array([[-1], [m.n_cells[0]]])) == jnp.array([-1, -2]) ) cell_idxs, faces_types = m._face_fidx_to_face_midx_and_face_type(jnp.arange(m.n_faces)) expected_cells_idx = jnp.array([[0], [1], [2], [3]]) expected_faces_types = jnp.array([0] * 4) assert jnp.all(cell_idxs == expected_cells_idx) assert jnp.all(faces_types == expected_faces_types) assert all( i == m.cell_midx_and_face_type_to_face_fidx(cell_idxs[i], faces_types[i]) for i in range(m.n_faces) ) a, b = m._face_fidx_to_neighbors_fidx(jnp.arange(m.n_faces)) assert all( i == m.get_face_index_from_two_neighbors(a[i], b[i]) for i in range(m.n_faces) ) x = jnp.array([0.1]) cell_idx, flat_idx = m.find_cell_index(x) print("x: ", x) print("multi index :", cell_idx) print("flat index :", flat_idx) print("\n\n@@@@@@@@@@@@@@@@ test 2d mesh (2,3) @@@@@@@@@@@@@@@") m = Mesh( dim=2, n_cells=(2, 3), ref_quad=UnitSquareTensorized(dim=2, order=3), mapping=Mapping([mapping_id]), ) try: m.get_face_index_from_two_neighbors(1, 1) except AssertionError: pass try: m.get_face_index_from_two_neighbors(-1, -3) except AssertionError: pass try: m.get_face_index_from_two_neighbors(-1, 3) except AssertionError: pass try: m.get_face_index_from_two_neighbors(0, 5) except AssertionError: pass assert m.get_face_index_from_two_neighbors(0, -1) == 0 assert m.get_face_index_from_two_neighbors(3, -2) == 6 assert m.get_face_index_from_two_neighbors(0, -3) == 9 assert m.get_face_index_from_two_neighbors(2, -4) == 12 assert m.get_face_index_from_two_neighbors(0, 1) == 10 assert m.get_face_index_from_two_neighbors(0, 3) == 3 for i in range(m.n_cells_total): assert i == m._midx_to_fidx(m._fidx_to_midx(i)) assert jnp.all(m.n_i_faces == jnp.array([9, 8])) assert m.n_faces == 17 assert jnp.all(m.faces_stride == jnp.array([0, 9, 17])) assert jnp.all( m._midx_to_fidx_s( jnp.array([[-1, 0], [m.n_cells[0], 0], [0, -1], [0, m.n_cells[1]]]) ) == jnp.array([-1, -2, -3, -4]) ) cell_idxs, faces_types = m._face_fidx_to_face_midx_and_face_type(jnp.arange(m.n_faces)) expected_cells_idx = jnp.array( [ [0, 0], [0, 1], [0, 2], [1, 0], [1, 1], [1, 2], [2, 0], [2, 1], [2, 2], [0, 0], [0, 1], [0, 2], [0, 3], [1, 0], [1, 1], [1, 2], [1, 3], ] ) expected_faces_types = jnp.array([0] * 9 + [1] * 8) assert jnp.all(cell_idxs == expected_cells_idx) assert jnp.all(faces_types == expected_faces_types) assert all( i == m.cell_midx_and_face_type_to_face_fidx(cell_idxs[i], faces_types[i]) for i in range(m.n_faces) ) a, b = m._face_fidx_to_neighbors_fidx(jnp.arange(m.n_faces)) assert all( i == m.get_face_index_from_two_neighbors(a[i], b[i]) for i in range(m.n_faces) ) x = jnp.array([0.1, 0.7]) cell_idx, flat_idx = m.find_cell_index(x) print("x: ", x) print("multi index :", cell_idx) print("flat index :", flat_idx) # plot_mesh(m) print("\n\n@@@@@@@@@@@@@@@@ test 2d mesh (4,3) @@@@@@@@@@@@@@@") m = Mesh( dim=2, n_cells=(4, 3), ref_quad=UnitSquareTensorized(dim=2, order=3), mapping=Mapping([mapping_id]), ) for i in range(m.n_cells_total): assert i == m._midx_to_fidx(m._fidx_to_midx(i)) assert jnp.all( m._midx_to_fidx_s( jnp.array([[-1, 0], [m.n_cells[0], 0], [0, -1], [0, m.n_cells[1]]]) ) == jnp.array([-1, -2, -3, -4]) ) assert jnp.all(m.n_i_faces == jnp.array([15, 16])) assert m.n_faces == 31 assert jnp.all(m.faces_stride == jnp.array([0, 15, 31])) cell_idxs, faces_types = m._face_fidx_to_face_midx_and_face_type(jnp.arange(m.n_faces)) assert all( i == m.cell_midx_and_face_type_to_face_fidx(cell_idxs[i], faces_types[i]) for i in range(m.n_faces) ) a, b = m._face_fidx_to_neighbors_fidx(jnp.arange(31)) assert all( i == m.get_face_index_from_two_neighbors(a[i], b[i]) for i in range(m.n_faces) ) x = jnp.array([0.1, 0.7]) cell_idx, flat_idx = m.find_cell_index(x) print("x: ", x) print("multi index :", cell_idx) print("flat index :", flat_idx) print("\n\n@@@@@@@@@@@@@@@@ test 3d mesh (4,3,5) @@@@@@@@@@@@@@@") m = Mesh( dim=3, n_cells=(4, 3, 5), ref_quad=UnitSquareTensorized(dim=3, order=3), mapping=Mapping([mapping_id]), ) # plot_mesh(m, show_indices_faces=False, show_indices_cells=True) # plot_mesh(m, direction=0) for i in range(m.n_cells_total): print(f"cell {i} <-> midx {m._fidx_to_midx(i)}") assert i == m._midx_to_fidx(m._fidx_to_midx(i)) assert jnp.all( m._midx_to_fidx_s( jnp.array( [ [-1, 0, 0], [m.n_cells[0], 0, 0], [0, -1, 0], [0, m.n_cells[1], 0], [0, 0, -1], [0, 0, m.n_cells[2]], ] ) ) == jnp.array([-1, -2, -3, -4, -5, -6]) ) assert jnp.all(m.n_i_faces == jnp.array([75, 80, 72])) assert m.n_faces == 227 assert jnp.all(m.faces_stride == jnp.array([0, 75, 155, 227])) cell_idxs, faces_types = m._face_fidx_to_face_midx_and_face_type(jnp.arange(m.n_faces)) assert all( i == m.cell_midx_and_face_type_to_face_fidx(cell_idxs[i], faces_types[i]) for i in range(m.n_faces) ) a, b = m._face_fidx_to_neighbors_fidx(jnp.arange(m.n_faces)) assert all( i == m.get_face_index_from_two_neighbors(a[i], b[i]) for i in range(m.n_faces) ) x = jnp.array([0.1, 0.7, 0.5]) cell_idx, flat_idx = m.find_cell_index(x) print("x: ", x) print("multi index :", cell_idx) print("flat index :", flat_idx) print("\n\n@@@@@@@@@@@@@@@@ test 3d mesh (4,3,5,2) @@@@@@@@@@@@@@@") m = Mesh( dim=4, n_cells=(4, 3, 5, 2), ref_quad=UnitSquareTensorized(dim=4, order=3), mapping=Mapping([mapping_id]), ) for i in range(m.n_cells_total): assert i == m._midx_to_fidx(m._fidx_to_midx(i)) assert jnp.all( m._midx_to_fidx_s( jnp.array( [ [-1, 0, 0, 0], [m.n_cells[0], 0, 0, 0], [0, -1, 0, 0], [0, m.n_cells[1], 0, 0], [0, 0, -1, 0], [0, 0, m.n_cells[2], 0], [0, 0, 0, -1], [0, 0, 0, m.n_cells[3]], ] ) ) == jnp.array([-1, -2, -3, -4, -5, -6, -7, -8]) ) assert jnp.all(m.n_i_faces == jnp.array([150, 160, 144, 180])) assert m.n_faces == 634 assert jnp.all(m.faces_stride == jnp.array([0, 150, 310, 454, 634])) cell_idxs, faces_types = m._face_fidx_to_face_midx_and_face_type(jnp.arange(m.n_faces)) assert all( i == m.cell_midx_and_face_type_to_face_fidx(cell_idxs[i], faces_types[i]) for i in range(m.n_faces) ) def get_neighbors(indices: jnp.ndarray): return m._face_fidx_to_neighbors_fidx(indices) get_neighbors_jitted = jax.jit(get_neighbors) start = timeit.default_timer() a, b = get_neighbors_jitted(jnp.arange(m.n_faces)) print("a[0]: ", a[0]) end = timeit.default_timer() print("time in get neighbors first time: ", end - start) start = timeit.default_timer() a, b = get_neighbors_jitted(jnp.arange(m.n_faces)) print("a[0]: ", a[0]) end = timeit.default_timer() print("time in get neighbors second time: ", end - start) assert all( i == m.get_face_index_from_two_neighbors(a[i], b[i]) for i in range(m.n_faces) ) x = jnp.array([0.1, 0.7, 0.5, 0.7]) cell_idx, flat_idx = m.find_cell_index(x) print("x: ", x) print("multi index :", cell_idx) print("flat index :", flat_idx)