ngsolve-webgpu 0.0.1__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.
Files changed (37) hide show
  1. ngsolve_webgpu/__init__.py +0 -0
  2. ngsolve_webgpu/_version.py +21 -0
  3. ngsolve_webgpu/animate.py +76 -0
  4. ngsolve_webgpu/cf.py +369 -0
  5. ngsolve_webgpu/clipping.py +155 -0
  6. ngsolve_webgpu/geometry.py +376 -0
  7. ngsolve_webgpu/isosurface.py +101 -0
  8. ngsolve_webgpu/jupyter.py +168 -0
  9. ngsolve_webgpu/lic.py +115 -0
  10. ngsolve_webgpu/mesh.py +424 -0
  11. ngsolve_webgpu/shaders/clipping/common.wgsl +120 -0
  12. ngsolve_webgpu/shaders/clipping/compute.wgsl +30 -0
  13. ngsolve_webgpu/shaders/clipping/render.wgsl +37 -0
  14. ngsolve_webgpu/shaders/compute.wgsl +59 -0
  15. ngsolve_webgpu/shaders/elements3d.wgsl +83 -0
  16. ngsolve_webgpu/shaders/eval/common.wgsl +14 -0
  17. ngsolve_webgpu/shaders/eval/seg.wgsl +30 -0
  18. ngsolve_webgpu/shaders/eval/tet.wgsl +61 -0
  19. ngsolve_webgpu/shaders/eval/trig.wgsl +121 -0
  20. ngsolve_webgpu/shaders/eval.wgsl +0 -0
  21. ngsolve_webgpu/shaders/geo_edge.wgsl +55 -0
  22. ngsolve_webgpu/shaders/geo_face.wgsl +47 -0
  23. ngsolve_webgpu/shaders/geo_vertex.wgsl +71 -0
  24. ngsolve_webgpu/shaders/isosurface/compute.wgsl +63 -0
  25. ngsolve_webgpu/shaders/isosurface/negative_clipping.wgsl +13 -0
  26. ngsolve_webgpu/shaders/isosurface/negative_surface.wgsl +15 -0
  27. ngsolve_webgpu/shaders/isosurface/render.wgsl +50 -0
  28. ngsolve_webgpu/shaders/line_integral_convolution.wgsl +100 -0
  29. ngsolve_webgpu/shaders/mesh.wgsl +296 -0
  30. ngsolve_webgpu/shaders/numbers.wgsl +40 -0
  31. ngsolve_webgpu/shaders/shader.wgsl +111 -0
  32. ngsolve_webgpu/shaders/uniforms.wgsl +93 -0
  33. ngsolve_webgpu-0.0.1.dist-info/METADATA +13 -0
  34. ngsolve_webgpu-0.0.1.dist-info/RECORD +37 -0
  35. ngsolve_webgpu-0.0.1.dist-info/WHEEL +5 -0
  36. ngsolve_webgpu-0.0.1.dist-info/licenses/LICENSE +504 -0
  37. ngsolve_webgpu-0.0.1.dist-info/top_level.txt +1 -0
File without changes
@@ -0,0 +1,21 @@
1
+ # file generated by setuptools-scm
2
+ # don't change, don't track in version control
3
+
4
+ __all__ = ["__version__", "__version_tuple__", "version", "version_tuple"]
5
+
6
+ TYPE_CHECKING = False
7
+ if TYPE_CHECKING:
8
+ from typing import Tuple
9
+ from typing import Union
10
+
11
+ VERSION_TUPLE = Tuple[Union[int, str], ...]
12
+ else:
13
+ VERSION_TUPLE = object
14
+
15
+ version: str
16
+ __version__: str
17
+ __version_tuple__: VERSION_TUPLE
18
+ version_tuple: VERSION_TUPLE
19
+
20
+ __version__ = version = '0.0.1'
21
+ __version_tuple__ = version_tuple = (0, 0, 1)
@@ -0,0 +1,76 @@
1
+ from webgpu.render_object import RenderObject
2
+ import ngsolve as ngs
3
+
4
+
5
+ class Animation(RenderObject):
6
+ def __init__(self, child):
7
+ super().__init__()
8
+ self.child = child
9
+ self.data = child.data
10
+ self.time_index = -1
11
+ self.max_time = -1
12
+ self.gfs = set()
13
+ self.parameters = dict()
14
+ f = self.data.cf
15
+ self.crawl_function(f)
16
+ # initial solution
17
+ self.add_time(initial=True)
18
+ self.store = True
19
+
20
+ def update(self, timestamp):
21
+ self.child.options = self.options
22
+ self.child.update(timestamp)
23
+
24
+ def get_bounding_box(self):
25
+ return self.child.get_bounding_box()
26
+
27
+ def crawl_function(self, f):
28
+ if f is None:
29
+ return
30
+ if isinstance(f, ngs.GridFunction):
31
+ self.gfs.add(f)
32
+ elif isinstance(f, ngs.Parameter) or isinstance(f, ngs.ParameterC):
33
+ self.parameters[f] = []
34
+ else:
35
+ for c in f.data["childs"]:
36
+ self.crawl_function(c)
37
+
38
+ def add_time(self, initial=False):
39
+ self.max_time += 1
40
+ self.time_index = self.max_time
41
+ for gf in self.gfs:
42
+ gf.AddMultiDimComponent(gf.vec)
43
+ for par, vals in self.parameters.items():
44
+ vals.append(par.Get())
45
+ if not initial:
46
+ self.slider.max(self.max_time)
47
+ # set value triggers set_time_index
48
+ self.slider.setValue(self.time_index)
49
+
50
+ def redraw(self, timestamp: float | None = None):
51
+ if self.store:
52
+ self.add_time()
53
+ else:
54
+ self.child.redraw(timestamp)
55
+
56
+ def render(self, encoder):
57
+ self.child.render(encoder)
58
+
59
+ def add_options_to_gui(self, gui):
60
+ self.slider = gui.slider(
61
+ 0,
62
+ self.set_time_index,
63
+ min=0,
64
+ max=0,
65
+ step=1,
66
+ label="animate",
67
+ )
68
+ self.child.add_options_to_gui(gui)
69
+
70
+ def set_time_index(self, time_index):
71
+ self.time_index = time_index
72
+ for gf in self.gfs:
73
+ gf.vec.data = gf.vecs[time_index + 1]
74
+ for p, vals in self.parameters.items():
75
+ p.Set(vals[time_index])
76
+ self.child.redraw()
ngsolve_webgpu/cf.py ADDED
@@ -0,0 +1,369 @@
1
+ import math
2
+
3
+ import ngsolve as ngs
4
+ import ngsolve.webgui
5
+ import numpy as np
6
+ from webgpu.clipping import Clipping
7
+ from webgpu.colormap import Colormap
8
+ from webgpu.render_object import RenderObject
9
+ from webgpu.utils import (
10
+ BufferBinding,
11
+ UniformBinding,
12
+ buffer_from_array,
13
+ read_shader_file,
14
+ )
15
+ from webgpu.vectors import BaseVectorRenderObject, VectorRenderer
16
+ from webgpu.webgpu_api import Buffer
17
+
18
+ from .mesh import Binding as MeshBinding, Mesh2dElementsRenderer
19
+ from .mesh import ElType, MeshData
20
+
21
+
22
+ class Binding:
23
+ FUNCTION_VALUES_2D = 10
24
+ COMPONENT = 55
25
+
26
+ _intrules_3d = {}
27
+
28
+
29
+ def get_3d_intrules(order):
30
+ if order in _intrules_3d:
31
+ return _intrules_3d[order]
32
+ ref_pts = [
33
+ [(order - i - j - k) / order, k / order, j / order]
34
+ for i in range(order + 1)
35
+ for j in range(order + 1 - i)
36
+ for k in range(order + 1 - i - j)
37
+ ]
38
+ p1_tets = {ngs.ET.TET: [[(1, 0, 0), (0, 1, 0), (0, 0, 1), (0, 0, 0)]]}
39
+ p1_tets[ngs.ET.PYRAMID] = [
40
+ [(1, 0, 0), (0, 1, 0), (0, 0, 1), (0, 0, 0)],
41
+ [(1, 0, 0), (0, 1, 0), (0, 0, 1), (1, 1, 0)],
42
+ ]
43
+ p1_tets[ngs.ET.PRISM] = [
44
+ [(1, 0, 0), (0, 1, 0), (0, 0, 1), (0, 0, 0)],
45
+ [(0, 0, 1), (0, 1, 0), (0, 1, 1), (1, 0, 0)],
46
+ [(1, 0, 1), (0, 1, 1), (1, 0, 0), (0, 0, 1)],
47
+ ]
48
+ p1_tets[ngs.ET.HEX] = [
49
+ [(1, 0, 0), (0, 1, 0), (0, 0, 1), (0, 0, 0)],
50
+ [(0, 1, 1), (1, 1, 1), (1, 1, 0), (1, 0, 1)],
51
+ [(1, 0, 1), (0, 1, 1), (1, 0, 0), (0, 0, 1)],
52
+ [(0, 1, 1), (1, 1, 0), (0, 1, 0), (1, 0, 0)],
53
+ [(0, 0, 1), (0, 1, 0), (0, 1, 1), (1, 0, 0)],
54
+ [(1, 0, 1), (1, 1, 0), (0, 1, 1), (1, 0, 0)],
55
+ ]
56
+ rules = {}
57
+ if order > 1:
58
+ ho_tets = {}
59
+ for eltype in p1_tets:
60
+ for tet in p1_tets[eltype]:
61
+ ho_tets[eltype] = []
62
+ for lam in ref_pts:
63
+ lami = [*lam, 1 - sum(lam)]
64
+ ho_tets[eltype].append(
65
+ [sum([lami[j] * tet[j][i] for j in range(4)]) for i in range(3)]
66
+ )
67
+ rules[eltype] = ngs.IntegrationRule(ho_tets[eltype])
68
+ else:
69
+ for eltype in p1_tets:
70
+ rules[eltype] = ngs.IntegrationRule(sum(p1_tets[eltype], []))
71
+ _intrules_3d[order] = rules
72
+ return rules
73
+
74
+
75
+ def _get_bernstein_matrix_trig(n, intrule):
76
+ """Create inverse vandermonde matrix for the Bernstein basis functions on a triangle of degree n and given integration points"""
77
+ ndtrig = int((n + 1) * (n + 2) / 2)
78
+
79
+ mat = ngs.Matrix(ndtrig, ndtrig)
80
+ fac_n = math.factorial(n)
81
+ for row, ip in enumerate(intrule):
82
+ col = 0
83
+ x = 1.0 - ip.point[0] - ip.point[1]
84
+ y = ip.point[1]
85
+ z = 1.0 - x - y
86
+ for i in range(n + 1):
87
+ factor = fac_n / math.factorial(i) * x**i
88
+ for j in range(n + 1 - i):
89
+ k = n - i - j
90
+ factor2 = 1.0 / (math.factorial(j) * math.factorial(k))
91
+ mat[row, col] = factor * factor2 * y**j * z**k
92
+ col += 1
93
+ return mat
94
+
95
+
96
+ def evaluate_cf(cf, mesh, order):
97
+ """Evaluate a coefficient function on a mesh and returns the values as a flat array, ready to copy to the GPU as storage buffer.
98
+ The first two entries are the function dimension and the polynomial order of the stored values.
99
+ """
100
+ comps = cf.dim
101
+ int_points = ngsolve.webgui._make_trig(order)
102
+ intrule = ngs.IntegrationRule(
103
+ int_points,
104
+ [
105
+ 0,
106
+ ]
107
+ * len(int_points),
108
+ )
109
+ ibmat = _get_bernstein_matrix_trig(order, intrule).I
110
+
111
+ ndof = ibmat.h
112
+
113
+ if isinstance(mesh, ngs.Region):
114
+ if mesh.VB() == ngs.VOL and mesh.mesh.dim == 3:
115
+ region = mesh.Boundaries()
116
+ else:
117
+ region = mesh
118
+ else:
119
+ region = mesh.Materials(".*")
120
+ if mesh.dim == 3:
121
+ region = mesh.Boundaries(".*")
122
+ pts = region.mesh.MapToAllElements(
123
+ {ngs.ET.TRIG: intrule, ngs.ET.QUAD: intrule}, region
124
+ )
125
+ pmat = cf(pts)
126
+ minval, maxval = (
127
+ (min(pmat.reshape(-1)), max(pmat.reshape(-1))) if len(pmat) else (0, 1)
128
+ )
129
+ pmat = pmat.reshape(-1, ndof, comps)
130
+
131
+ values = np.zeros((ndof, pmat.shape[0], comps), dtype=np.float32)
132
+ for i in range(comps):
133
+ ngsmat = ngs.Matrix(pmat[:, :, i].transpose())
134
+ values[:, :, i] = ibmat * ngsmat
135
+
136
+ values = values.transpose((1, 0, 2)).flatten()
137
+ ret = np.concatenate(([np.float32(cf.dim), np.float32(order)], values.reshape(-1)))
138
+ # print("ret = ", ret)
139
+ return ret, minval, maxval
140
+
141
+
142
+ class FunctionData:
143
+ mesh_data: MeshData
144
+ data_2d: np.ndarray | None = None
145
+ data_3d: np.ndarray | None = None
146
+ gpu_2d: Buffer | None = None
147
+ gpu_3d: Buffer | None = None
148
+ cf: ngs.CoefficientFunction
149
+ order: int
150
+ order_3d: int
151
+ _timestamp: float = -1
152
+ minval: float = 1e99
153
+ maxval: float = -1e99
154
+
155
+ def __init__(
156
+ self,
157
+ mesh_data: MeshData,
158
+ cf: ngs.CoefficientFunction,
159
+ order: int,
160
+ order3d: int = -1,
161
+ ):
162
+ self.mesh_data = mesh_data
163
+ self.cf = cf
164
+ self.order = order
165
+ self.order_3d = order if order3d == -1 else order3d
166
+ self.need_3d = False
167
+
168
+ def update(self, timestamp: float):
169
+ if self._timestamp == timestamp:
170
+ return
171
+ self._timestamp = timestamp
172
+ if self.need_3d:
173
+ self.mesh_data.need_3d = True
174
+ self.mesh_data.update(timestamp)
175
+ self._create_data()
176
+
177
+ def _create_data(self):
178
+ self.gpu_2d = None
179
+ self.gpu_3d = None
180
+ self.data_2d, self.minval, self.maxval = evaluate_cf(
181
+ self.cf, self.mesh_data.ngs_mesh, self.order
182
+ )
183
+ if self.need_3d:
184
+ self.data_3d, minval, maxval = self.evaluate_3d(
185
+ self.cf, self.mesh_data.ngs_mesh, self.order_3d
186
+ )
187
+ self.minval = min(self.minval, minval)
188
+ self.maxval = max(self.maxval, maxval)
189
+
190
+ def get_buffers(self):
191
+ buffers = self.mesh_data.get_buffers().copy()
192
+ if self.gpu_2d is None:
193
+ self.gpu_2d = buffer_from_array(self.data_2d)
194
+ if self.data_3d is not None:
195
+ self.gpu_3d = buffer_from_array(self.data_3d)
196
+ buffers["data_2d"] = self.gpu_2d
197
+ if self.gpu_3d is not None:
198
+ buffers["data_3d"] = self.gpu_3d
199
+ self.data_2d = None
200
+ self.data_3d = None
201
+ return buffers
202
+
203
+ def get_bounding_box(self):
204
+ return self.mesh_data.get_bounding_box()
205
+
206
+ def evaluate_3d(self, cf, region, order):
207
+ intrules = get_3d_intrules(order)
208
+ if not isinstance(region, ngs.Region):
209
+ region = region.Materials(".*")
210
+ pts = region.mesh.MapToAllElements(intrules, region)
211
+ V_inv = vandermonde_3d(order).T
212
+ vals = cf(pts).reshape(-1, len(intrules[ngs.ET.TET])).dot(V_inv)
213
+ vmin, vmax = vals.min(), vals.max()
214
+ ret = np.concatenate(
215
+ ([np.float32(cf.dim), np.float32(order)], vals.reshape(-1)),
216
+ dtype=np.float32,
217
+ )
218
+ return ret, vmin, vmax
219
+
220
+
221
+ _vandermonde_mats = {}
222
+
223
+
224
+ def vandermonde_3d(order):
225
+ if order in _vandermonde_mats:
226
+ return _vandermonde_mats[order]
227
+ basis_indices = [
228
+ (order - i - j - k, k, j, i)
229
+ for i in range(order + 1)
230
+ for j in range(order + 1 - i)
231
+ for k in range(order + 1 - i - j)
232
+ ]
233
+ n = len(basis_indices)
234
+ V = np.zeros((n, n))
235
+ for r, (i, j, k, l) in enumerate(basis_indices):
236
+ for c, (a, b, c2, d) in enumerate(basis_indices):
237
+ multinom_coef = math.factorial(order) / (
238
+ math.factorial(a)
239
+ * math.factorial(b)
240
+ * math.factorial(c2)
241
+ * math.factorial(d)
242
+ )
243
+ V[r, c] = (
244
+ multinom_coef
245
+ * (i / order) ** a
246
+ * (j / order) ** b
247
+ * (k / order) ** c2
248
+ * (l / order) ** d
249
+ )
250
+ _vandermonde_mats[order] = np.linalg.inv(V)
251
+ return _vandermonde_mats[order]
252
+
253
+
254
+ class CFRenderer(Mesh2dElementsRenderer):
255
+ """Use "vertices", "index" and "trig_function_values" buffers to render a mesh"""
256
+ fragment_entry_point = "fragmentTrig"
257
+
258
+ def __init__(self, data: FunctionData, component=0, label="CFRenderer"):
259
+ super().__init__(data=data.mesh_data, label=label)
260
+ self.data = data
261
+ self.colormap = Colormap()
262
+ self.component = component
263
+
264
+ def update(self, timestamp):
265
+ if timestamp == self._timestamp:
266
+ return
267
+ self._timestamp = timestamp
268
+ self.data.update(timestamp)
269
+ self._buffers = self.data.get_buffers()
270
+ self.colormap.options = self.options
271
+
272
+ self.curvature_subdivision = self.data.mesh_data.curvature_subdivision
273
+ self.n_vertices = 3 * self.curvature_subdivision**2
274
+ if self.colormap.autoupdate:
275
+ self.colormap.set_min_max(
276
+ self.data.minval, self.data.maxval, set_autoupdate=False
277
+ )
278
+ self.colormap.update(timestamp)
279
+ self.clipping.update(timestamp)
280
+ self.n_instances = self.data.mesh_data.num_elements[ElType.TRIG]
281
+ self.component_buffer = buffer_from_array(np.array([self.component], np.int32))
282
+ self.create_render_pipeline()
283
+
284
+ def get_bounding_box(self):
285
+ return self.data.get_bounding_box()
286
+
287
+ def add_options_to_gui(self, gui):
288
+ if self.data.cf.dim > 1:
289
+ options = {"Norm": 0}
290
+ for d in range(self.data.cf.dim):
291
+ options[str(d)] = d + 1
292
+ gui.dropdown(func=self.change_cf_dim, label="Component", values=options)
293
+
294
+ def change_cf_dim(self, value):
295
+ self.component = value
296
+ self.component_buffer = buffer_from_array(np.array([self.component], np.int32))
297
+ self.options.render_function()
298
+
299
+ def get_shader_code(self):
300
+ shader_code = ""
301
+
302
+ for file_name in [
303
+ "eval.wgsl",
304
+ "mesh.wgsl",
305
+ "shader.wgsl",
306
+ "uniforms.wgsl",
307
+ ]:
308
+ shader_code += read_shader_file(file_name, __file__)
309
+
310
+ shader_code += self.colormap.get_shader_code()
311
+ shader_code += self.clipping.get_shader_code()
312
+ shader_code += self.options.camera.get_shader_code()
313
+ shader_code += self.options.light.get_shader_code()
314
+ return shader_code
315
+
316
+ def get_bindings(self):
317
+ return [*super().get_bindings(),
318
+ *self.colormap.get_bindings(),
319
+ BufferBinding(Binding.FUNCTION_VALUES_2D, self._buffers["data_2d"]),
320
+ BufferBinding(Binding.COMPONENT, self.component_buffer)]
321
+
322
+ class VectorCFRenderer(VectorRenderer):
323
+ def __init__(
324
+ self, cf: ngs.CoefficientFunction, mesh: ngs.Mesh, grid_size=20, size=None,
325
+ ):
326
+ # calling super-super class to not create points and vectors
327
+ BaseVectorRenderObject.__init__(self)
328
+ self.cf = cf
329
+ self.mesh = mesh
330
+ # this somehow segfaults in pyodide?
331
+ self.grid_size = grid_size
332
+ self.size = size
333
+
334
+ def redraw(self, timestamp=None):
335
+ super().redraw(
336
+ timestamp=timestamp, cf=self.cf, mesh=self.mesh, grid_size=self.grid_size
337
+ )
338
+
339
+ def update(self, timestamp: float):
340
+ if self._timestamp == timestamp:
341
+ return
342
+ bb = self.mesh.ngmesh.bounding_box
343
+ self.bounding_box = np.array(
344
+ [[bb[0][0], bb[0][1], bb[0][2]], [bb[1][0], bb[1][1], bb[1][2]]]
345
+ )
346
+ vs = np.linspace(
347
+ self.bounding_box[0][0],
348
+ self.bounding_box[1][0],
349
+ self.grid_size + 1,
350
+ endpoint=False,
351
+ )[1:]
352
+ points = np.meshgrid(vs, vs)
353
+ xvals = points[0].flatten()
354
+ yvals = points[1].flatten()
355
+ self.size = self.size or 1 / 60 * np.linalg.norm(
356
+ self.bounding_box[1] - self.bounding_box[0]
357
+ )
358
+ mpts_ = self.mesh(xvals, yvals, 0.0)
359
+ pts, mpts = [], []
360
+ for i in range(len(xvals)):
361
+ if mpts_[i]["nr"] != -1:
362
+ mpts.append(mpts_[i])
363
+ pts.append([xvals[i], yvals[i], 0.0])
364
+ self.points = np.array(pts, dtype=np.float32).reshape(-1)
365
+ values = self.cf(mpts)
366
+ self.vectors = np.array(
367
+ [values[:, 0], values[:, 1], np.zeros_like(values[:, 0])], dtype=np.float32
368
+ ).T.reshape(-1)
369
+ super().update(timestamp)
@@ -0,0 +1,155 @@
1
+ from webgpu import create_bind_group, read_shader_file
2
+ from webgpu.utils import buffer_from_array, uniform_from_array
3
+ from webgpu.clipping import Clipping
4
+ from webgpu.colormap import Colormap
5
+ from webgpu.render_object import RenderObject
6
+ from webgpu.utils import BufferBinding, UniformBinding, ReadBuffer
7
+
8
+ from webgpu.webgpu_api import *
9
+
10
+ import numpy as np
11
+
12
+ from .cf import FunctionData
13
+
14
+ from .mesh import Mesh3dElementsRenderObject, ElType
15
+ from .mesh import Binding as MeshBinding
16
+
17
+
18
+ class VolumeCF(Mesh3dElementsRenderObject):
19
+ fragment_entry_point: str = "cf_fragment_main"
20
+
21
+ def __init__(self, data: FunctionData):
22
+ super().__init__(data=data.mesh_data)
23
+ self.data = data
24
+ self.data.need_3d = True
25
+ self.colormap = Colormap()
26
+
27
+ def update(self, timestamp):
28
+ if self._timestamp == timestamp:
29
+ return
30
+ self.colormap.options = self.options
31
+ self.colormap.update(timestamp)
32
+ super().update(timestamp)
33
+
34
+ def get_bindings(self):
35
+ return super().get_bindings() + [
36
+ BufferBinding(10, self._buffers["data_3d"]),
37
+ *self.colormap.get_bindings(),
38
+ ]
39
+
40
+ def get_shader_code(self):
41
+ eval_code = read_shader_file("eval.wgsl", __file__)
42
+ return super().get_shader_code() + self.colormap.get_shader_code() + eval_code
43
+
44
+
45
+ class ClippingCF(RenderObject):
46
+ compute_shader = "clipping/compute.wgsl"
47
+ n_vertices = 3
48
+ subdivision = 0
49
+
50
+ def __init__(self, data: FunctionData):
51
+ super().__init__()
52
+ self.clipping = Clipping()
53
+ self.colormap = Colormap()
54
+ self.clipping.callbacks.append(self.build_clip_plane)
55
+ self.data = data
56
+ self.data.need_3d = True
57
+
58
+ def update(self, timestamp):
59
+ if timestamp == self._timestamp:
60
+ return
61
+ self._timestamp = timestamp
62
+ self.data.update(timestamp)
63
+ self.clipping.update(timestamp)
64
+ self.colormap.options = self.options
65
+ self.colormap.update(timestamp)
66
+ self._buffers = self.data.get_buffers()
67
+ self.build_clip_plane()
68
+
69
+ def get_bounding_box(self):
70
+ return self.data.get_bounding_box()
71
+
72
+ def get_shader_code(self, compute=False):
73
+ shader_code = ""
74
+ shader_code += self.clipping.get_shader_code()
75
+ shader_code += self.options.camera.get_shader_code()
76
+ shader_code += read_shader_file("clipping/common.wgsl", __file__)
77
+ shader_code += read_shader_file("eval/common.wgsl", __file__)
78
+ shader_code += read_shader_file("eval/tet.wgsl", __file__)
79
+ if compute:
80
+ shader_code += read_shader_file(self.compute_shader, __file__)
81
+ else:
82
+ shader_code += read_shader_file("clipping/render.wgsl", __file__)
83
+ shader_code += self.colormap.get_shader_code()
84
+ shader_code += self.options.light.get_shader_code()
85
+ return shader_code
86
+
87
+ def get_bindings(self, compute=False):
88
+ bindings = [
89
+ *self.options.camera.get_bindings(),
90
+ BufferBinding(MeshBinding.VERTICES, self._buffers["vertices"]),
91
+ UniformBinding(22, self.n_tets),
92
+ UniformBinding(23, self.only_count),
93
+ BufferBinding(MeshBinding.TET, self._buffers[ElType.TET]),
94
+ BufferBinding(13, self._buffers["data_3d"]),
95
+ *self.clipping.get_bindings(),
96
+ ]
97
+ if compute:
98
+ bindings += [
99
+ BufferBinding(
100
+ 21,
101
+ self.trig_counter,
102
+ read_only=False,
103
+ visibility=ShaderStage.COMPUTE,
104
+ ),
105
+ BufferBinding(24, self.cut_trigs, read_only=False),
106
+ ]
107
+ else:
108
+ bindings += [
109
+ *self.colormap.get_bindings(),
110
+ *self.options.light.get_bindings(),
111
+ BufferBinding(24, self.cut_trigs),
112
+ ]
113
+ return bindings
114
+
115
+ def build_clip_plane(self):
116
+ for count in [True, False]:
117
+ encoder = self.device.createCommandEncoder("build_clip_plane")
118
+ ntets = self.data.mesh_data.num_elements[ElType.TET] * 4**self.subdivision
119
+ self.trig_counter = buffer_from_array(
120
+ np.array([0], dtype=np.uint32),
121
+ usage=BufferUsage.STORAGE | BufferUsage.COPY_DST | BufferUsage.COPY_SRC,
122
+ )
123
+ self.n_tets = uniform_from_array(np.array([ntets], dtype=np.uint32))
124
+ self.only_count = uniform_from_array(np.array([count], dtype=np.uint32))
125
+ if count:
126
+ self.cut_trigs = buffer_from_array(
127
+ np.array([0.0] * 64, dtype=np.float32)
128
+ )
129
+ else:
130
+ self.cut_trigs = self.device.createBuffer(
131
+ size=64 * self.n_instances, usage=BufferUsage.STORAGE
132
+ )
133
+ layout, group = create_bind_group(
134
+ self.device, self.get_bindings(compute=True), label="create_clip_plane"
135
+ )
136
+ shader_module = self.device.createShaderModule(
137
+ code=self.get_shader_code(compute=True)
138
+ )
139
+ pipeline = self.device.createComputePipeline(
140
+ self.device.createPipelineLayout([layout]),
141
+ label="create_clip_plane",
142
+ compute=ComputeState(module=shader_module, entryPoint="main"),
143
+ )
144
+ compute_pass = encoder.beginComputePass(label="build_clip_plane")
145
+ compute_pass.setPipeline(pipeline)
146
+ compute_pass.setBindGroup(0, group)
147
+ compute_pass.dispatchWorkgroups(1024)
148
+ compute_pass.end()
149
+ if count:
150
+ read = ReadBuffer(self.trig_counter, encoder)
151
+ self.device.queue.submit([encoder.finish()])
152
+ if count:
153
+ array = read.get_array(dtype=np.uint32)
154
+ self.n_instances = int(array[0])
155
+ self.create_render_pipeline()