-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathscaling_jax_sharding.py
More file actions
100 lines (73 loc) · 3.33 KB
/
Copy pathscaling_jax_sharding.py
File metadata and controls
100 lines (73 loc) · 3.33 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
import os
import jax
# os.environ["CUDA_VISIBLE_DEVICES"] = "1,2,3,4,5,6,7,8"
from jax.experimental.shard_map import shard_map
from tinyfluids.jax_tinyfluids.time_integration import _halo_exchange, _time_integration_inner, time_integration
from tinyfluids.jax_tinyfluids.sharding_helpers import pad, unpad
from tinyfluids.jax_tinyfluids.fluid import DENSITY_INDEX, PRESSURE_INDEX, X_AXIS, Y_AXIS, Z_AXIS, VAR_AXIS
import timeit
import jax.numpy as jnp
from jax.sharding import PartitionSpec as P, NamedSharding
import matplotlib.pyplot as plt
def plot_results(final_state):
num_cells = final_state.shape[1]
fig, axs = plt.subplots(1, 2, figsize=(12, 6))
axs[0].imshow(final_state[DENSITY_INDEX, :, :, num_cells // 2], extent = [0, 1, 0, 1])
axs[0].set_title("Density")
axs[1].imshow(final_state[PRESSURE_INDEX, :, :, num_cells // 2], extent = [0, 1, 0, 1])
axs[1].set_title("Pressure")
plt.savefig("figures/check_{:d}.png".format(num_cells))
def setup_ics(num_cells, num_injection_cells=2):
grid_spacing = 1 / (num_cells - 1)
rho = jnp.ones((num_cells, num_cells, num_cells)) * 0.125
u_x = jnp.zeros((num_cells, num_cells, num_cells))
u_y = jnp.zeros((num_cells, num_cells, num_cells))
u_z = jnp.zeros((num_cells, num_cells, num_cells))
p = jnp.ones((num_cells, num_cells, num_cells)) * 0.1
center = num_cells // 2
injection_slice = slice(center - num_injection_cells, center + num_injection_cells)
rho = rho.at[injection_slice, injection_slice, injection_slice].set(1.0)
p = p.at[injection_slice, injection_slice, injection_slice].set(1.0)
primitive_state = jnp.stack([rho, u_x, u_y, u_z, p], axis = 0)
return primitive_state, grid_spacing
num_cells = 256
num_injection_cells = num_cells // 16
primitive_state, grid_spacing = setup_ics(num_cells, num_injection_cells)
t_final = 0.2
gamma = 5/3
shard = True
shard_mapped = True
# TODO: do outer boarders better
if shard:
split = (1, 2, 2, 2)
sharding_mesh = jax.make_mesh(split, (VAR_AXIS, X_AXIS, Y_AXIS, Z_AXIS))
sharding = jax.NamedSharding(sharding_mesh, P(VAR_AXIS, X_AXIS, Y_AXIS, Z_AXIS))
primitive_state = jax.device_put(primitive_state, sharding)
padding = ((0, 0), (1, 1), (1, 1), (1, 1))
if shard_mapped:
primitive_state = pad(primitive_state, padding, sharding)
print("started first run")
# Execute once for compilation and warmup
if shard_mapped:
final_state, num_iterations = time_integration(primitive_state, grid_spacing, t_final, gamma, shard_mapped, padding, split)
else:
final_state, num_iterations = _time_integration_inner(primitive_state, grid_spacing, t_final, gamma, shard_mapped)
final_state.block_until_ready()
print("finished first run")
if shard_mapped:
final_state = unpad(final_state, padding, sharding)
plot_results(final_state)
def time_execution():
if shard_mapped:
final_state, _ = time_integration(primitive_state, grid_spacing, t_final, gamma, shard_mapped, padding, split)
else:
final_state, _ = _time_integration_inner(primitive_state, grid_spacing, t_final, gamma, shard_mapped)
final_state.block_until_ready()
# Measure execution time
times = timeit.repeat(
time_execution,
repeat = 3, # More repeats for better statistics
number = 1 # Number of calls per measurement
)
print(times)
print(f"Execution time: {min(times)} seconds")