MR-AIV reveals in vivo brain-wide fluid flow with physics-informed AI.
The 1 match
- [1] § MATERIALS AND METHODS › Magnetic resonance artificial intelligence velocimetry ↔ src/instant_aiv/models/NNpp.py, lines 201–233 · score 0.52 · activation function, weight normalization, connections, network, layers, models
Paper
Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC
The paper is loaded when this pane is shown.
The authors' code
Python · 824 lines · 31 KB · no license · 1 match
- # Libraries
- import numpy as np
- from jax import jit, vmap, grad
- import jax.numpy as jnp
- from jax.nn import sigmoid
- from typing import Tuple
- from typing import Tuple, List, Dict, Sequence
- import h5py
- from instant_aiv.models.metrics import *
- from instant_aiv.manage.dataloader import *
- import optax
- #Initialization
- from typing import Tuple
- def glorot_normal(in_dim: int, out_dim: int) -> jnp.ndarray:
- glorot_stddev = np.sqrt(2.0 / (in_dim + out_dim))
- return jnp.array(np.random.normal(loc=0.0, scale=glorot_stddev, size=(in_dim, out_dim)))
- def init_params(layers: List[int], initialization_type: str = 'xavier',Network_type: str='mlp',degree: int =5,Use_ResNet: bool =False) -> dict:
- def init_adaptive_params():
- F = 0.1 * jnp.ones(3 * len(layers) - 1)
- A = 0.1 * jnp.ones(3 * len(layers) - 1)
- return [{"a0": A[3*i], "a1": A[3*i + 1], "a2": A[3*i + 2],
- "f0": F[3*i], "f1": F[3*i + 1], "f2": F[3*i + 2]}
- for i in range(len(layers) - 1)]
- #Define Models:
- def init_layer_mlp(in_dim, out_dim):
- if initialization_type == 'xavier':
- W = glorot_normal(in_dim, out_dim)
- elif initialization_type == 'normal':
- W = jnp.array(np.random.normal(size=(in_dim, out_dim)))
- b = jnp.zeros(out_dim)
- g = jnp.ones(out_dim)
- return {"W": W, "b": b, "g": g}
- def init_layer_ResNet(in_dim, out_dim):
- if initialization_type == 'xavier':
- W1 = glorot_normal(in_dim, out_dim)
- W2 = glorot_normal(out_dim, out_dim)
- W3 = glorot_normal(out_dim, out_dim)
- elif initialization_type == 'normal':
- W1 = jnp.array(np.random.normal(size=(in_dim, out_dim)))
- W2 = jnp.array(np.random.normal(size=(out_dim, out_dim)))
- W3 = jnp.array(np.random.normal(size=(out_dim, out_dim)))
- b1 = jnp.zeros(out_dim)
- b2 = jnp.zeros(out_dim)
- b3 = jnp.zeros(out_dim)
- g1 = jnp.ones(out_dim)
- g2 = jnp.ones(out_dim)
- g3 = jnp.ones(out_dim)
- alpha=0.0
- return {"W": W1, "b": b1, "g": g1,
- "W2": W2, "b2": b2, "g2": g2,
- 'alpha':alpha}
- def init_layer_kan(in_dim, out_dim,degree=degree):
- std=1 / (in_dim * (degree + 1))
- W =jnp.array(np.random.normal(loc=0.0, scale=std, size=(in_dim, out_dim,degree+1)))
- b = jnp.zeros(out_dim)
- g = jnp.ones(out_dim)
- return {"W": W, "b": b, "g": g}
- def init_layer_kan_ResNet(in_dim, out_dim,degree=degree):
- std=1 / (in_dim * (degree + 1))
- W1 =jnp.array(np.random.normal(loc=0.0, scale=std, size=(in_dim, out_dim,degree+1)))
- W2 =jnp.array(np.random.normal(loc=0.0, scale=std, size=(in_dim, out_dim,degree+1)))
- b1 = jnp.zeros(out_dim)
- b2 = jnp.zeros(out_dim)
- g1 = jnp.ones(out_dim)
- g2 = jnp.ones(out_dim)
- alpha=0.0
- return {"W": W1, "b": b1, "g": g1,
- "W2": W2, "b2": b2, "g2": g2,
- 'alpha':alpha}
- #Select model
- if Network_type.lower()=='mlp':
- if Use_ResNet:
- init_layer_params=init_layer_ResNet
- else:
- init_layer_params=init_layer_mlp
- elif Network_type.lower()[:3]=='kan':
- if Use_ResNet:
- init_layer_params=init_layer_kan_ResNet
- else:
- init_layer_params=init_layer_kan
- else:
- print(f'Error: {Network_type.lower()} is not a valid option. The available options are:mlp and kan.')
- print(f'Initializing:{Network_type} parameters.')
- params = [init_layer_params(layers[i], layers[i + 1]) for i in range(len(layers) - 1)]
- U1, b1, g1 = glorot_normal(layers[0], layers[1]), jnp.zeros(layers[1]), jnp.ones(layers[1])
- U2, b2, g2 = glorot_normal(layers[0], layers[1]), jnp.zeros(layers[1]), jnp.ones(layers[1])
- mMLP_params = [{"U1": U1, "b1": b1, "g1": g1, "U2": U2, "b2": b2, "g2": g2}]
- return {
- 'params': params,
- 'AdaptiveAF': init_adaptive_params(),
- 'mMLP': mMLP_params
- }
- def init_params_res(layers: List[int], initialization_type: str = 'xavier',Network_type: str='mlp',degree: int =5,Use_ResNet: bool =False) -> dict:
- def init_adaptive_params():
- F = 0.1 * jnp.ones(3 * len(layers) - 1)
- A = 0.1 * jnp.ones(3 * len(layers) - 1)
- return [{"a0": A[3*i], "a1": A[3*i + 1], "a2": A[3*i + 2],
- "f0": F[3*i], "f1": F[3*i + 1], "f2": F[3*i + 2]}
- for i in range(len(layers) - 1)]
- def init_layer_params(in_dim, out_dim):
- if initialization_type == 'xavier':
- W1 = glorot_normal(in_dim, out_dim)
- W2 = glorot_normal(out_dim, out_dim)
- W3 = glorot_normal(out_dim, out_dim)
- elif initialization_type == 'normal':
- W1 = jnp.array(np.random.normal(size=(in_dim, out_dim)))
- W2 = jnp.array(np.random.normal(size=(out_dim, out_dim)))
- W3 = jnp.array(np.random.normal(size=(out_dim, out_dim)))
- b1 = jnp.zeros(out_dim)
- b2 = jnp.zeros(out_dim)
- b3 = jnp.zeros(out_dim)
- g1 = jnp.ones(out_dim)
- g2 = jnp.ones(out_dim)
- g3 = jnp.ones(out_dim)
- alpha=0.0
- return {"W": W1, "b": b1, "g": g1,
- "W2": W2, "b2": b2, "g2": g2,
- 'alpha':alpha}
- params = [init_layer_params(layers[i], layers[i + 1]) for i in range(len(layers) - 1)]
- U1, b1, g1 = glorot_normal(layers[0], layers[1]), jnp.zeros(layers[1]), jnp.ones(layers[1])
- U2, b2, g2 = glorot_normal(layers[0], layers[1]), jnp.zeros(layers[1]), jnp.ones(layers[1])
- mMLP_params = [{"U1": U1, "b1": b1, "g1": g1, "U2": U2, "b2": b2, "g2": g2}]
- return {
- 'params': params,
- 'AdaptiveAF': init_adaptive_params(),
- 'mMLP': mMLP_params
- }
- def init_params_dict(layer_dict, initialization,Use_ResNet=False,Network_type='mlp',degree=5):
- print(f'You selected: Network {Network_type} with degree(if KAN) {degree}, initialization {initialization},Use_ResNet {Use_ResNet}')
- if Network_type.lower()=='mlp':
- if Use_ResNet:
- init_function=init_params_res
- else:
- init_function=init_params
- elif Network_type[:3].lower()=='kan':
- init_function=init_params
- initialized_params = {}
- for key, layer_structure in layer_dict.items():
- # Initialize parameters for each key
- params = init_function(layer_structure,
- initialization_type=initialization.lower(),
- Network_type=Network_type,
- degree=degree,
- Use_ResNet=Use_ResNet)
- # Store in the dictionary
- initialized_params[key] = params
- return initialized_params
- def FCN(params, X_in, M1, M2, activation_fn, norm_fn):
- """
- Fully Connected Network (FCN) with a given normalization and activation function.
- Parameters:
- - params: Dictionary containing parameters of the neural network layers.
- - X_in: Input tensor to the network.
- - M1, M2: Parameters for normalization function.
- - activation_fn: Callable activation function.
- - norm_fn: Callable normalization function.
- Returns:
- - Output tensor from the network.
- """
- params_N = params["params"]
- inputs = norm_fn(X_in, M1, M2)
- for layer in params_N[:-1]:
- outputs = activation_fn(jnp.dot(inputs, layer["W"]) + layer["b"])
- inputs = outputs
- W = params_N[-1]["W"]
- b = params_N[-1]["b"]
- outputs = jnp.dot(inputs, W) + b
- return outputs
- def FCN_WN(params, X_in, M1, M2, activation, norm_fn):
- """
- Fully Connected Network (FCN) with Weight Normalization.
- Parameters:
- - params: Dictionary containing parameters of the neural network layers.
- - X_in: Input tensor to the network.
- - M1, M2: Parameters for normalization function.
- - activation: Callable activation function.
- - norm_fn: Callable normalization function.
- Returns:
- - Output tensor from the network.
- """
- # Normalize the input
- H = norm_fn(X_in, M1, M2)
- # Iterate through the layers
- for i, layer_params in enumerate(params["params"]):
- W, b, g = layer_params["W"], layer_params["b"], layer_params["g"]
- # Weight normalization
- V = W / jnp.linalg.norm(W, axis=0, keepdims=True)
- # Linear transformation
- H = g * jnp.matmul(H, V) + b
- # Apply activation function for all layers except the last one
- if i != len(params["params"]) - 1:
- H = activation(H)
- return H
- def FCN_WN_MMLP(params, X_in, M1, M2, activation, norm_fn):
- """
- Fully Connected Network (FCN) with Weight Normalization and Modified MLP (MMLP) transformations.
- Parameters:
- - params: Dictionary containing parameters of the neural network layers and MMLP parameters.
- - X_in: Input tensor to the network.
- - M1, M2: Parameters for normalization function.
- - activation: Callable activation function.
- - norm_fn: Callable normalization function.
- Returns:
- - Output tensor from the network.
- """
- # Normalize the input
- H = norm_fn(X_in, M1, M2)
- # Unpack MMLP parameters and apply weight normalization
- mMLP_params = params["mMLP"][0]
- U1, U2, b1, b2, g1, g2 = (mMLP_params[key] for key in ["U1", "U2", "b1", "b2", "g1", "g2"])
- U1_norm, U2_norm = U1 / jnp.linalg.norm(U1, axis=0, keepdims=True), U2 / jnp.linalg.norm(U2, axis=0, keepdims=True)
- # Calculate U and V transformations
- U = activation(g1 * jnp.dot(H, U1_norm) + b1)
- V = activation(g2 * jnp.dot(H, U2_norm) + b2)
- # Iterate through the layers
- for idx, layer in enumerate(params["params"][:-1]):
- W, b, g = layer["W"], layer["b"], layer["g"]
- # Apply weight normalization
- W_norm = W / jnp.linalg.norm(W, axis=0, keepdims=True)
- # Compute activations and apply MMLP combination step
- H = activation(g * jnp.dot(H, W_norm) + b)
- H = jnp.multiply(H, U) + jnp.multiply(1 - H, V)
- # Process the last layer
- W, b, g = params["params"][-1]["W"], params["params"][-1]["b"], params["params"][-1]["g"]
- W_norm = W / jnp.linalg.norm(W, axis=0, keepdims=True)
- H = g * jnp.dot(H, W_norm) + b
- return H
- def FCN_MMLP(params, X_in, M1, M2, activation, norm_fn):
- """
- Fully Connected Network (FCN) with Modified MLP (MMLP) transformations.
- Parameters:
- - params: Dictionary containing parameters of the neural network layers and MMLP parameters.
- - X_in: Input tensor to the network.
- - M1, M2: Parameters for normalization function.
- - activation: Callable activation function.
- - norm_fn: Callable normalization function.
- Returns:
- - Output tensor from the network.
- """
- # Normalize the input
- inputs = norm_fn(X_in, M1, M2)
- # Unpack MMLP parameters
- mMLP_params = params["mMLP"][0]
- U1, U2, b1, b2 = mMLP_params["U1"], mMLP_params["U2"], mMLP_params["b1"], mMLP_params["b2"]
- # Calculate U and V transformations
- U = activation(jnp.dot(inputs, U1) + b1)
- V = activation(jnp.dot(inputs, U2) + b2)
- # Iterate through all layers except the last
- for layer in params["params"][:-1]:
- W, b = layer["W"], layer["b"]
- # Compute activations
- act_values = activation(jnp.dot(inputs, W) + b)
- # MMLP combination step
- inputs = jnp.multiply(act_values, U) + jnp.multiply(1 - act_values, V)
- # Compute output from the last layer
- W, b = params["params"][-1]["W"], params["params"][-1]["b"]
- outputs = jnp.dot(inputs, W) + b
- return outputs
- # Save Resuls
- def save_list(Loss,path,name='loss-'):
- filename=path+name+".npy"
- np.save(filename, np.array(Loss))
- def save_MLP_params(params: List[Dict[str, np.ndarray]], save_path,WN=False,Mod_MLP=False):
- with h5py.File(save_path, "w") as f:
- for layer_idx, layer_params in enumerate(params):
- layer_group = f.create_group(f"Layer_{layer_idx/2.0:.2f}")
- if WN:
- if Mod_MLP:
- W, b, g= layer_params.values()
- layer_group.create_dataset("W", shape=W.shape, dtype=np.float32, data=W)
- layer_group.create_dataset("b", shape=b.shape, dtype=np.float32, data=b)
- layer_group.create_dataset("g", shape=g.shape, dtype=np.float32, data=g)
- else:
- W, b, g= layer_params.values()
- layer_group.create_dataset("W", shape=W.shape, dtype=np.float32, data=W)
- layer_group.create_dataset("b", shape=b.shape, dtype=np.float32, data=b)
- layer_group.create_dataset("g", shape=g.shape, dtype=np.float32, data=g)
- else:
- W, b = layer_params.values()
- layer_group.create_dataset("W", shape=W.shape, dtype=np.float32, data=W)
- layer_group.create_dataset("b", shape=b.shape, dtype=np.float32, data=b)
- def read_params(filename,WN=False):
- data = h5py.File(filename, 'r')
- recover_params=[]
- for layer in data.keys() :
- if WN:
- stored={'W':data[layer]['W'][:],'b': data[layer]['b'][:],'g': data[layer]['g'][:]}
- else:
- stored={'W':data[layer]['W'][:],'b': data[layer]['b'][:]}
- recover_params.append(stored)
- return recover_params
- def select_model(WN=False, Mod_MLP=False,Use_ResNet=False,Adaptive=False,Light=False,Network_type='mlp',degree=5):
- if Network_type.lower()=='mlp':
- model_map = {
- (True, True,False, False, False): FCN_WN_MMLP,
- (True, False,False, False, False): FCN_WN,
- (False, True,False, False, False): FCN_MMLP,
- (False, False,False, False, False): FCN,
- (False, False,True, False, False): ResNet,
- (False, False,True, False, True): ResNet_light,
- (True, False,True, False, False): WN_ResNet,
- (True, False,True, False, True): WN_ResNet_light,
- (True, False,True, True, False): WN_ResNet_adaptive,
- }
- return model_map[(WN, Mod_MLP,Use_ResNet,Adaptive,Light)]
- elif Network_type.lower()=='kan':
- if degree==3:
- return KAN_Net3
- elif degree==5:
- if Use_ResNet:
- return KAN5_ResNet
- else:
- return KAN_Net5
- elif degree==7:
- return KAN_Net7
- elif degree==9:
- return KAN_Net9
- elif degree==11:
- return KAN_Net11
- else:
- return KAN_Net
- elif Network_type.lower()=='kan_theta':
- return KAN_Net_theta
- def initialize_optimizer(lr0, decay_rate, lrf, decay_step, T_e,optimizer_type='Adam',weight_decay=1e-5):
- print('Optimizer',optimizer_type.lower())
- if optimizer_type.lower()=='adam':
- if decay_rate == 0 or lrf == lr0:
- print('No decay')
- return optax.adam(lr0), decay_step
- else:
- if decay_step == 0:
- decay_step = T_e * np.log(decay_rate) / np.log(lrf / lr0)
- print(f'The decay step will be {decay_step}')
- return optax.adam(optax.exponential_decay(lr0, decay_step, decay_rate,)),decay_step
- elif optimizer_type.lower()=='adamw':
- print('Weight decay:',weight_decay)
- if decay_rate == 0 or lrf == lr0:
- print('No decay')
- return optax.adamw(learning_rate=lr0, weight_decay=weight_decay), decay_step
- else:
- if decay_step == 0:
- decay_step = T_e * np.log(decay_rate) / np.log(lrf / lr0)
- print(f'The decay step will be {decay_step}')
- # Use adamw with the specified learning rate schedule
- return optax.adamw(optax.exponential_decay(lr0, decay_step, decay_rate), weight_decay=weight_decay), decay_step
- elif optimizer_type.lower()=='lion':
- if decay_rate == 0 or lrf == lr0:
- weight_decay=weight_decay*3
- print('No decay')
- return optax.lion(learning_rate=lr0, weight_decay=weight_decay), decay_step
- else:
- if decay_step == 0:
- weight_decay=weight_decay*3
- decay_step = T_e * np.log(decay_rate) / np.log(lrf / lr0)
- print(f'The decay step will be {decay_step}')
- # Use adamw with the specified learning rate schedule
- return optax.lion(optax.exponential_decay(lr0, decay_step, decay_rate), weight_decay=weight_decay), decay_step
- def load_params_dict(result_path, dataset_name, layer_dict, initialization, type='Test',Use_ResNet=False):
- loaded_params = {}
- for key in layer_dict.keys():
- # Construct the file path
- file_path = f"{result_path}{dataset_name}-{type}_params_{key}.h5"
- # Read parameters from the file
- raw_params = read_all_params(file_path)[0]
- # Initialize a test parameter set for getting lengths
- test_params = init_params(layer_dict[key], initialization_type=initialization.lower())
- params_length = len(test_params['params'])
- params_length_AF = len(test_params['AdaptiveAF'])
- # Reconstruct the parameters
- reconstructed_params = reconstruct_params(raw_params, params_length, params_length_AF,Use_ResNet)
- # Store in the dictionary
- loaded_params[key] = reconstructed_params
- return loaded_params
- def save_params_dict(params,result_path,dataset_name,type='Test',Use_ResNet=False):
- # Assuming 'result_path' and 'dataset_name' are defined elsewhere in your code
- for key, params in params.items():
- # Construct the file name for each set of parameters
- output_path = f"{result_path}{dataset_name}-{type}_params_{key}.h5"
- # Extract arrays from params
- params_to_save = [extract_arrays_from_params(params,Use_ResNet)]
- # Save the parameters
- save_all_params(params_to_save, output_path)
- print(f'Params {key} have been saved!')
- def transfer_params(params_s,params_t,levels=[0,1,2]):
- for level in levels:
- params_t['params'][level]=params_s['params'][level]
- return params_t
- def L2_regularization(params,beta=0.001):
- Loss_L2=0
- for param in params:
- Loss_L2=jnp.mean(param['W']**2)+Loss_L2
- return beta*Loss_L2
- def L1_regularization(params,beta=0.001):
- Loss_L2=0
- for param in params:
- Loss_L2=jnp.abs(jnp.mean(param['W']))+Loss_L2
- return beta*Loss_L2
- def L2_regularization_ResNet(params,beta=0.001):
- Loss_L2=0
- for param in params:
- Loss_L2=jnp.sum(param['W']**2)+jnp.sum(param['W2']**2)+Loss_L2
- return beta*Loss_L2
- # RESNET ARCHITECTURES
- def ResNet(params, X_in, M1, M2, activation_fn, norm_fn):
- def linear_layer(H,layer_params,W_key="W",b_key="b"):
- W, b= layer_params[W_key], layer_params[b_key]
- H= jnp.dot(H,W) + b
- return H
- params_N = params["params"]
- inputs = norm_fn(X_in, M1, M2)
- layer_params=params_N[0]
- inputs=activation_fn(linear_layer(inputs,layer_params,W_key="W",b_key="b"))
- for ly in range(1,len(params_N[:-1])):
- layer_params=params_N[ly]
- g= activation_fn(linear_layer(inputs,layer_params,W_key="W",b_key="b"))
- h= linear_layer(g,layer_params,W_key="W2",b_key="b2")
- inputs= activation_fn(h+inputs)
- outputs = jnp.dot(inputs, params_N[-1]["W"]) + params_N[-1]["b"]
- return outputs
- def ResNet_light(params, X_in, M1, M2, activation_fn, norm_fn):
- def linear_layer(H,layer_params,W_key="W",b_key="b"):
- W, b= layer_params[W_key], layer_params[b_key]
- H= jnp.dot(H,W) + b
- return H
- params_N = params["params"]
- inputs = norm_fn(X_in, M1, M2)
- layer_params=params_N[0]
- inputs=activation_fn(linear_layer(inputs,layer_params,W_key="W",b_key="b"))
- inputs0=jnp.copy(inputs)
- for ly in range(1,len(params_N[:-2])):
- layer_params=params_N[ly]
- g= activation_fn(linear_layer(inputs,layer_params,W_key="W",b_key="b"))
- h= linear_layer(g,layer_params,W_key="W2",b_key="b2")
- inputs= activation_fn(h)
- layer_params=params_N[-2]
- inputs=activation_fn(linear_layer(inputs+inputs0,layer_params,W_key="W",b_key="b"))
- layer_params=params_N[-1]
- outputs=activation_fn(linear_layer(inputs,layer_params,W_key="W",b_key="b"))
- return outputs
- def WN_ResNet(params, X_in, M1, M2, activation, norm_fn):
- def linear_layer(H,layer_params,W_key="W",b_key="b",g_key="g"):
- W, b, g = layer_params[W_key], layer_params[b_key], layer_params[g_key]
- # Weight normalization
- V = W / jnp.linalg.norm(W, axis=0, keepdims=True)
- # Linear transformation
- H = g * jnp.matmul(H, V) + b
- return H
- params_N = params["params"]
- # Normalize the input
- H = norm_fn(X_in, M1, M2)
- # Encoder Layer
- H = activation(linear_layer(H,params_N[0],W_key="W",b_key="b",g_key="g"))
- # Iterate through the layers
- for i in range(1,len(params_N[:-1])):
- layer_params=params_N[i]
- F=activation(linear_layer(H,layer_params,W_key="W",b_key="b",g_key="g"))
- G=linear_layer(F,layer_params,W_key="W2",b_key="b2",g_key="g2")
- H=activation(G+H)
- H=linear_layer(H,params_N[-1],W_key="W",b_key="b",g_key="g")
- return H
- def WN_ResNet_adaptive(params, X_in, M1, M2, activation, norm_fn):
- def linear_layer(H,layer_params,W_key="W",b_key="b",g_key="g"):
- W, b, g = layer_params[W_key], layer_params[b_key], layer_params[g_key]
- # Weight normalization
- V = W / jnp.linalg.norm(W, axis=0, keepdims=True)
- # Linear transformation
- H = g * jnp.matmul(H, V) + b
- return H
- params_N = params["params"]
- # Normalize the input
- H = norm_fn(X_in, M1, M2)
- # Encoder Layer
- H = activation(linear_layer(H,params_N[0],W_key="W",b_key="b",g_key="g"))
- # Iterate through the layers
- for i in range(1,len(params_N[:-1])):
- layer_params=params_N[i]
- F=activation(linear_layer(H,layer_params,W_key="W",b_key="b",g_key="g"))
- G=linear_layer(F,layer_params,W_key="W2",b_key="b2",g_key="g2")
- H=activation(layer_params["alpha"]*G+(1-layer_params["alpha"])*H)
- H=linear_layer(H,params_N[-1],W_key="W",b_key="b",g_key="g")
- return H
- def WN_ResNet_light(params, X_in, M1, M2, activation, norm_fn):
- def linear_layer(H,layer_params,W_key="W",b_key="b",g_key="g"):
- W, b, g = layer_params[W_key], layer_params[b_key], layer_params[g_key]
- # Weight normalization
- V = W / jnp.linalg.norm(W, axis=0, keepdims=True)
- # Linear transformation
- H = g * jnp.matmul(H, V) + b
- return H
- params_N = params["params"]
- # Normalize the input
- H = norm_fn(X_in, M1, M2)
- # Encoder Layer
- H = activation(linear_layer(H,params_N[0],W_key="W",b_key="b",g_key="g"))
- H0=jnp.copy(H)
- # Iterate through the layers
- for i in range(1,len(params_N[:-2])):
- layer_params=params_N[i]
- F=activation(linear_layer(H,layer_params,W_key="W",b_key="b",g_key="g"))
- G=linear_layer(F,layer_params,W_key="W2",b_key="b2",g_key="g2")
- H=activation(G)
- H=activation(linear_layer(H+H0,params_N[-2],W_key="W",b_key="b",g_key="g"))
- H=linear_layer(H,params_N[-1],W_key="W",b_key="b",g_key="g")
- return H
- def KAN_Net3(params, X_in, M1, M2, activation, norm_fn):
- def Cheby_KAN_layer3(x,layer_params,expanded_arr=[]):
- # Read chebyshev coefficients:
- cheby_coeffs= layer_params["W"]
- inputdim=cheby_coeffs.shape[0]
- outdim=cheby_coeffs.shape[1]
- # Normalize
- x = activation(x)
- # Reshape
- x = x.reshape((-1, inputdim, 1))
- x=jnp.stack((T0(x),
- T1(x),
- T2(x),
- T3(x)),axis=2)# Discard dummy dimension
- # Compute the Chebyshev interpolation
- x = jnp.einsum("bid0,iod->bo", x, cheby_coeffs) # shape = (batch_size, output_dim)
- # Remove extra dimension
- x= jnp.reshape(x, (outdim,))
- return x
- #Define params
- params_N = params["params"]
- #Normalize inputs
- x = norm_fn(X_in, M1, M2)
- for ly in range(len(params_N)):
- layer_params=params_N[ly]
- x=Cheby_KAN_layer3(x,layer_params)
- return x
- def KAN_Net5(params, X_in, M1, M2, activation, norm_fn):
- def Cheby_KAN_layer5(x,layer_params,expanded_arr=[]):
- # Read chebyshev coefficients:
- cheby_coeffs= layer_params["W"]
- inputdim=cheby_coeffs.shape[0]
- outdim=cheby_coeffs.shape[1]
- # Normalize
- x = activation(x)
- # Reshape
- x = x.reshape((-1, inputdim, 1))
- x=jnp.stack((T0(x),
- T1(x),
- T2(x),
- T3(x),
- T4(x),
- T5(x)),axis=2)# Discard dummy dimension
- # Compute the Chebyshev interpolation
- x = jnp.einsum("bid0,iod->bo", x, cheby_coeffs) # shape = (batch_size, output_dim)
- # Remove extra dimension
- x= jnp.reshape(x, (outdim,))
- return x
- #Define params
- params_N = params["params"]
- #Normalize inputs
- x = norm_fn(X_in, M1, M2)
- for ly in range(len(params_N)):
- layer_params=params_N[ly]
- x=Cheby_KAN_layer5(x,layer_params)
- return x
- def KAN_Net7(params, X_in, M1, M2, activation, norm_fn):
- def Cheby_KAN_layer7(x,layer_params,expanded_arr=[]):
- # Read chebyshev coefficients:
- cheby_coeffs= layer_params["W"]
- inputdim=cheby_coeffs.shape[0]
- outdim=cheby_coeffs.shape[1]
- # Normalize
- x = activation(x)
- # Reshape
- x = x.reshape((-1, inputdim, 1))
- x=jnp.stack((T0(x),
- T1(x),
- T2(x),
- T3(x),
- T4(x),
- T5(x),
- T6(x),
- T7(x)),axis=2)# Discard dummy dimension
- # Compute the Chebyshev interpolation
- x = jnp.einsum("bid0,iod->bo", x, cheby_coeffs) # shape = (batch_size, output_dim)
- # Remove extra dimension
- x= jnp.reshape(x, (outdim,))
- return x
- #Define params
- params_N = params["params"]
- #Normalize inputs
- x = norm_fn(X_in, M1, M2)
- for ly in range(len(params_N)):
- layer_params=params_N[ly]
- x=Cheby_KAN_layer7(x,layer_params)
- return x
- def KAN_Net9(params, X_in, M1, M2, activation, norm_fn):
- def Cheby_KAN_layer9(x,layer_params,expanded_arr=[]):
- # Read chebyshev coefficients:
- cheby_coeffs= layer_params["W"]
- inputdim=cheby_coeffs.shape[0]
- outdim=cheby_coeffs.shape[1]
- # Normalize
- x = activation(x)
- # Reshape
- x = x.reshape((-1, inputdim, 1))
- x=jnp.stack((T0(x),
- T1(x),
- T2(x),
- T3(x),
- T4(x),
- T5(x),
- T6(x),
- T7(x),
- T8(x),
- T9(x)),axis=2)# Discard dummy dimension
- # Compute the Chebyshev interpolation
- x = jnp.einsum("bid0,iod->bo", x, cheby_coeffs) # shape = (batch_size, output_dim)
- # Remove extra dimension
- x= jnp.reshape(x, (outdim,))
- return x
- #Define params
- params_N = params["params"]
- #Normalize inputs
- x = norm_fn(X_in, M1, M2)
- for ly in range(len(params_N)):
- layer_params=params_N[ly]
- x=Cheby_KAN_layer9(x,layer_params)
- return x
- def KAN_Net11(params, X_in, M1, M2, activation, norm_fn):
- def Cheby_KAN_layer11(x,layer_params,expanded_arr=[]):
- # Read chebyshev coefficients:
- cheby_coeffs= layer_params["W"]
- inputdim=cheby_coeffs.shape[0]
- outdim=cheby_coeffs.shape[1]
- # Normalize
- x = activation(x)
- # Reshape
- x = x.reshape((-1, inputdim, 1))
- x=jnp.stack((T0(x),
- T1(x),
- T2(x),
- T3(x),
- T4(x),
- T5(x),
- T6(x),
- T7(x),
- T8(x),
- T9(x),
- T10(x),
- T11(x)),axis=2)# Discard dummy dimension
- # Compute the Chebyshev interpolation
- x = jnp.einsum("bid0,iod->bo", x, cheby_coeffs) # shape = (batch_size, output_dim)
- # Remove extra dimension
- x= jnp.reshape(x, (outdim,))
- return x
- #Define params
- params_N = params["params"]
- #Normalize inputs
- x = norm_fn(X_in, M1, M2)
- for ly in range(len(params_N)):
- layer_params=params_N[ly]
- x=Cheby_KAN_layer11(x,layer_params)
- return x
- # WN KAN
- def WN_KAN_Net5(params, X_in, M1, M2, activation, norm_fn):
- def Cheby_KAN_layer5(x,layer_params,expanded_arr=[]):
- # Read chebyshev coefficients and other params:
- cheby_coeffs= layer_params["W"]
- g= layer_params["g"]
- b= layer_params["b"]
- # Normalize
- cheby_coeffs= cheby_coeffs / jnp.linalg.norm(cheby_coeffs, axis=0, keepdims=True)
- inputdim=cheby_coeffs.shape[0]
- outdim=cheby_coeffs.shape[1]
- # Normalize
- x = activation(x)
- # Reshape
- x = x.reshape((-1, inputdim, 1))
- x=jnp.stack((T0(x),
- T1(x),
- T2(x),
- T3(x),
- T4(x),
- T5(x)),axis=2)# Discard dummy dimension
- # Compute the Chebyshev interpolation
- x = jnp.einsum("bid0,iod->bo", x, cheby_coeffs) # shape = (batch_size, output_dim)
- # Remove extra dimension
- x= jnp.reshape(x, (outdim,))
- return g * x+b
- #Define params
- params_N = params["params"]
- #Normalize inputs
- x = norm_fn(X_in, M1, M2)
- for ly in range(len(params_N)):
- layer_params=params_N[ly]
- x=Cheby_KAN_layer5(x,layer_params)
- return x
- def KAN5_ResNet(params, X_in, M1, M2, activation, norm_fn):
- def Cheby_KAN_layer5(x,layer_params,W_key):
- # Read chebyshev coefficients:
- cheby_coeffs= layer_params[W_key]
- inputdim=cheby_coeffs.shape[0]
- outdim=cheby_coeffs.shape[1]
- # Normalize
- x = activation(x)
- # Reshape
- x = x.reshape((-1, inputdim, 1))
- x=jnp.stack((T0(x),
- T1(x),
- T2(x),
- T3(x),
- T4(x),
- T5(x)),axis=2)# Discard dummy dimension
- # Compute the Chebyshev interpolation
- x = jnp.einsum("bid0,iod->bo", x, cheby_coeffs) # shape = (batch_size, output_dim)
- # Remove extra dimension
- x= jnp.reshape(x, (outdim,))
- return x
- params_N = params["params"]
- inputs = norm_fn(X_in, M1, M2)
- layer_params=params_N[0]
- inputs=Cheby_KAN_layer5(inputs,layer_params,W_key="W")
- for ly in range(1,len(params_N[:-1])):
- layer_params=params_N[ly]
- g= Cheby_KAN_layer5(inputs,layer_params,W_key="W")
- h= Cheby_KAN_layer5(g,layer_params,W_key="W2")
- inputs= h+inputs
- layer_params=params_N[-1]
- inputs=Cheby_KAN_layer5(inputs,layer_params,W_key="W")
- return inputs
NNpp.py at commit f242e08, no license · at the source
Overview
- Division of Applied Mathematics, Brown University, Providence, RI 02912, USA
- Department of Mechanical Engineering, University of Rochester, Rochester, NY 14627, USA
- School of Engineering, Brown University, Providence, RI 02912, USA
- Center for Translational Neuromedicine, University of Copenhagen, Copenhagen 2200, Denmark
Abstract
The circulation of cerebrospinal and interstitial fluid plays a vital role in clearing metabolic waste from the brain, and its disruption has been linked to neurological disorders. However, directly measuring brain-wide fluid transport, especially in the deep brain, has remained elusive. Here, we introduce magnetic resonance artificial intelligence velocimetry (MR-AIV), a framework featuring a specialized physics-informed architecture and optimization method that reconstructs three-dimensional fluid velocity fields from dynamic contrast-enhanced magnetic resonance imaging (DCE-MRI). MR-AIV unveils brain-wide velocity maps while providing estimates of tissue permeability and pressure fields, quantities inaccessible to other methods. Applied to the brain, MR-AIV reveals a functional landscape of interstitial and perivascular flow, quantitatively distinguishing slow diffusion-driven transport [∼0.1 micrometers per second (μm/
Reproduced under the paper's license (CC BY), from the paper cited above.
Repository
Its files are read in the Code ↔ Paper reader above, with 1 match between paragraphs and lines of code.
jdtoscano94/MR-AIV
f242e08e8c4f1fb82f395e363b258d4259600cba, 31 May 2026Availability: 1 check, the latest on 28 September 2026: the link answers
- 28 September 2026: the link answers
15 files
- Brain/
Permeability/ , Python, 963 linesInit_K.py - Brain/
Real/ , Python, 2,030 linesInit_c_and_p.py - Brain/
Real/ , Python, 3,155 linesTrain.py - Brain/
Real/ , Jupyter, 1,460 linesVisualizations.ipynb - Brain/
Synthetic/ , Python, 2,252 linesInit_c_and_p.py - Brain/
Synthetic/ , Python, 4,007 linesTrain.py - Brain/
Synthetic/ , Jupyter, 1,529 linesVisualizations.ipynb - src/
instant_aiv/ , Python, 1 line__init__.py - src/
instant_aiv/ , Python, 1 linemanage/ __init__.py - src/
instant_aiv/ , Python, 904 linesmanage/ dataloader.py - src/
instant_aiv/ , Python, 123 linesmanage/ plots.py - src/
instant_aiv/ , Python, 824 lines, 1 matchmodels/ NNpp.py - src/
instant_aiv/ , Python, 1 linemodels/ __init__.py - src/
instant_aiv/ , Python, 56 linesmodels/ metrics.py - README.md, Text, 412 lines
The paper's code and data availability statement is in the Data section.
Tracing map
Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.
What the map holds:
- 1 repository of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
- 14 scripts, each with its path and the digest of its content;
- 1 match between paragraphs of the paper and lines of the code (method lexical-v1);
- neither the text of the paper nor the code itself.
Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.
Data
Datasets cited
- zenodo:15345392, at Zenodo; found in “Data, code, and materials availability:”
Data, code, and materials availability
All data and code needed to evaluate and reproduce the results in the paper are present in the paper and/
Reproduced under the paper's license (CC BY), from the paper cited above.
Versions
The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.
Version 1, 28 September 2026: the first record
Recorded: type, language, journal, volume, issue, pages, dates, 8 authors, 6 MeSH terms, 4 funders, 69 references.
Cite
This paper
Toscano, J. D., Guo, Y., Wang, Z., Vaezi, M., Mori, Y., Karniadakis, G. E., Boster, K. A. S., & Kelley, D. H. (2026). MR-AIV reveals in vivo brain-wide fluid flow with physics-informed AI. Science advances, 12(22), eaeb0404. https://
BibTeX
@article{toscano2026mr,
author = {Toscano, Juan Diego and Guo, Yisen and Wang, Zhibo and Vaezi, Mohammad and Mori, Yuki and Karniadakis, George Em and Boster, Kimberly A. S. and Kelley, Douglas H.},
title = {{MR-AIV reveals in vivo brain-wide fluid flow with physics-informed AI}},
journal = {Science advances},
year = {2026},
month = may,
volume = {12},
number = {22},
pages = {eaeb0404},
publisher = {American Association for the Advancement of Science},
issn = {2375-2548},
doi = {10.1126/
url = {https://
pmid = {42202031},
pmcid = {PMC13215179}
}
RIS
TY - JOUR
AU - Toscano, Juan Diego
AU - Guo, Yisen
AU - Wang, Zhibo
AU - Vaezi, Mohammad
AU - Mori, Yuki
AU - Karniadakis, George Em
AU - Boster, Kimberly A. S.
AU - Kelley, Douglas H.
TI - MR-AIV reveals in vivo brain-wide fluid flow with physics-informed AI
T2 - Science advances
J2 - Sci Adv
PY - 2026
DA - 2026/
VL - 12
IS - 22
SP - eaeb0404
SN - 2375-2548
PB - American Association for the Advancement of Science
DO - 10.1126/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.1126/
"type": "article-journal",
"title": "MR-AIV reveals in vivo brain-wide fluid flow with physics-informed AI",
"container-title": "Science advances",
"author": [
{
"family": "Toscano",
"given": "Juan Diego"
},
{
"family": "Guo",
"given": "Yisen"
},
{
"family": "Wang",
"given": "Zhibo"
},
{
"family": "Vaezi",
"given": "Mohammad"
},
{
"family": "Mori",
"given": "Yuki"
},
{
"family": "Karniadakis",
"given": "George Em"
},
{
"family": "Boster",
"given": "Kimberly A. S."
},
{
"family": "Kelley",
"given": "Douglas H."
}
],
"container-title-short":
"volume": "12",
"issue": "22",
"page": "eaeb0404",
"DOI": "10.1126/
"PMID": "42202031",
"PMCID": "PMC13215179",
"ISSN": "2375-2548",
"publisher": "American Association for the Advancement of Science",
"URL": "https://
"language": "en",
"issued": {
"date-parts": [
[
2026,
5,
27
]
]
}
}
The tracing map gets a citation of its own once an author has validated it and it has a DOI.
Similar papers
The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.
- [1] doi:10.1186/s12987-026-00782-w
- Quantifying cerebrospinal fluid flow in pial perivascular spaces of rats.Journal: Fluids and barriers of the CNSIn common: 11 references, author Douglas H. Kelley
- [2] doi:10.1073/pnas.2526239123 [code]
- Quantitative assessment of flow between cerebrospinal and interstitial fluid compartments in humans.Journal: Proceedings of the National Academy of Sciences of the United States of AmericaIn common: structural MRI / diffusion, 8 references
- [3] doi:10.1038/s41467-026-76306-9 [code]
- Non-invasive characterization of perivascular subarachnoid spaces.Journal: Nature communicationsIn common: scikit-learn, SciPy, Matplotlib, 1 other tool, structural MRI / diffusion, 5 references
- [4] doi:10.1073/pnas.2517059123 [code]
- Stretch and flow at the gliovascular interface: High-fidelity modeling of astrocyte endfeet.Journal: Proceedings of the National Academy of Sciences of the United States of AmericaIn common: pandas, SciPy, Matplotlib, 1 other tool, 5 references
- [5] doi:10.1002/hipo.70131 [code]
- Decoding Medial Entorhinal Cortical Dynamics Produces Planning-Like Alternations in Hippocampal theta Sequences.Journal: HippocampusIn common: JAX, OpenCV, h5py, 5 other tools, systems
- [6] doi:10.1038/s41593-026-02262-8 [code]
- Cheese3D enables sensitive detection and analysis of whole-face movement in mice.Journal: Nature neuroscienceIn common: JAX, OpenCV, h5py, 5 other tools, systems
- [7] doi:10.1093/brain/awag080 [code]
- Early glymphatic failure in AppNL-F knock-in mice is linked to parenchymal border macrophages loss.Journal: Brain : a journal of neurologyIn common: 6 references
- [8] doi:10.1002/mrm.70514 [code]
- Brain Clearance of Contrast Agent in Intravenous DCE-MRI Is Measurable and Cannot Be Explained by Clearance Across the BBB Alone.Journal: Magnetic resonance in medicineIn common: pandas, SciPy, Matplotlib, 1 other tool, structural MRI / diffusion, 4 references
- [9] doi:10.1177/0271678x261455452 [code]
- Meningeal CSF transport varies across parasagittal dura subregions with age in humans.Journal: Journal of cerebral blood flow and metabolism : official journal of the International Society of Cerebral Blood Flow and MetabolismIn common: structural MRI / diffusion, 5 references
- [10] doi:10.1002/alz.71745 [code]
- Regional astrocyte dysregulation and altered glymphatic-related markers in Alzheimer's disease frontal cortex.Journal: Alzheimer's & dementia : the journal of the Alzheimer's AssociationIn common: OpenCV, pandas, SciPy, 2 other tools, 3 references
Contribute
The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.
Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.
Claim this paper
Correct its record
Say what each link of this record is, remove the ones that are not the paper's, add the ones that are missing. The correction becomes a new version of the record, in its Versions section.
Validate its tracing map
You validate the map as this page shows it: 1 repository of the authors' code, each at its verified commit and with its license, 14 scripts, and 1 match between paragraphs and code (see the Code and Map sections). It then receives a DOI on Zenodo, with you (your ORCID iD) and OSCR as its creators; the code itself is not deposited.
The map's fingerprint: sha256:853b80a7bd0881cb…
Add the badge to its README
The badge links the code to this page. Copy one of these into the README of the paper's code: only you decide where it goes, and nothing is changed for you.
Markdown
[, paste the snippet at the top, then “Commit changes…” and, to review it first, “Create a new branch and start a pull request”. You open the pull request; OSCR asks for no permission.
Request its removal
To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).
Discussion, reproductions, activity
Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.
Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.
Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.
