Neural Probabilistic Circuits: Enabling Compositional and Interpretable Predictions Through Logical Reasoning.
The 8 matches · 1 of them tie a paragraph to a whole file, not to given lines: a weak match, whose lines are not tinted
- [1] § Appendix 3: Experimental Setup ↔ src/npc-models/header.py, lines 1–67 · score 0.76 · concept loss weight, weight decay, momentum, embedding, validation, seeds
- [2] § Preliminaries ↔ src/npc-models/pc.py, lines 113–144 · score 0.67 · weighted sum, weight updating, root node, sum node, children, joint
- [3] § Neural Probabilistic Circuits › Three-Stage Training Algorithm › Circuit Construction ↔ src/npc-models/pc.py, lines 388–404 · score 0.66 · root node, leaf nodes, product node, sum node, depth
- [4] § Appendix 3: Experimental Setup ↔ src/npc-models/train_npc.py, lines 199–337 · score 0.64 · weight decay, SGD, plateaus, momentum, validation, seeds
- [5] § Preliminaries ↔ src/npc-models/train_npc.py, lines 199–337 · score 0.62 · leaf nodes, product node, sum node, Probabilistic circuits, transformed, loss
- [6] § Neural Probabilistic Circuits › Three-Stage Training Algorithm › Circuit Construction ↔ src/npc-models/train_pc.py, lines 61–151 · score 0.60 · leaf nodes, product nodes, sum nodes, likelihood, circuit, training
- [7] § Experiments › Ablation Studies › Impact of Interventions ↔ scripts/npc-models/train_npc.bash, the whole file · a weak match · score 0.59 · trained NPC models, joint optimization, CelebA, MNIST, GTSRB, AwA2
- [8] § Preliminaries ↔ src/npc-models/train_pc.py, lines 61–151 · score 0.56 · leaf node, product node, sum node, marginal, circuits, joint
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 · 829 lines · 32 KB · CC-BY-NC-SA-4.0 · 2 matches
- """
- @file pc.py
- @author Simon Yu
- @date 02/07/2024
- @brief PC classes.
- """
- import abc
- import itertools
- import logger
- import numpy
- import os
- import torch
- import tqdm
- class PCLearningRateScheduler:
- def __init__(self, optimizer, factor = 0.8, patience = 2, threshold = 1e-4, cooldown = 2, min_learning_rate = 1e-6):
- self.cooldown = cooldown
- self.cooldown_counter = 0
- self.factor = factor
- self.metrics = None
- self.min_learning_rate = min_learning_rate
- self.optimizer = optimizer
- self.patience = patience
- self.patience_counter = 0
- self.threshold = threshold
- return
- @abc.abstractmethod
- def step(self, metrics):
- pass
- class LikelihoodPCLearningRateScheduler(PCLearningRateScheduler):
- def __init__(self, optimizer, factor = 0.8):
- super().__init__(optimizer, factor)
- return
- def step(self, metrics):
- if self.metrics is not None:
- if metrics < self.metrics:
- self.optimizer.learning_rate *= self.factor
- logger.log_info("Reducing PC learning rate to " + "{:e}".format(self.optimizer.learning_rate) + "...")
- self.metrics = metrics
- return
- class LossPCLearningRateScheduler(PCLearningRateScheduler):
- def __init__(self, optimizer, factor = 0.8, patience = 2, threshold = 1e-4, cooldown = 2, min_learning_rate = 1e-6):
- super().__init__(optimizer, factor, patience, threshold, cooldown, min_learning_rate)
- return
- def step(self, metrics):
- if self.metrics is None:
- self.metrics = metrics
- return
- if self.optimizer.learning_rate < self.min_learning_rate:
- self.metrics = metrics
- return
- if self.cooldown_counter > 0:
- self.cooldown_counter -= 1
- self.metrics = metrics
- return
- if abs(self.metrics - metrics) < self.threshold:
- if self.patience_counter < self.patience:
- self.patience_counter += 1
- self.metrics = metrics
- return
- self.optimizer.learning_rate *= self.factor
- logger.log_info("Reducing PC learning rate to " + "{:e}".format(self.optimizer.learning_rate) + "...")
- self.cooldown_counter = self.cooldown
- self.patience_counter = 0
- else:
- self.patience_counter = 0
- self.metrics = metrics
- return
- class PCOptimizer:
- def __init__(self, pc_joint, pc_marginal, device = torch.device("cuda"), learning_rate = 1e-1, prior_factor = 1e2, projection_epsilon = 1e-2):
- self.device = device
- self.learning_rate = learning_rate
- self.smoothing_epsilon = torch.finfo(torch.float).eps
- self.prior_factor = prior_factor
- self.projection_epsilon = projection_epsilon
- self.pc_joint = pc_joint
- self.pc_marginal = pc_marginal
- self.weights_prior = None
- return
- def set_weights_prior(self, weights_prior):
- self.weights_prior = weights_prior
- for i in range(len(self.weights_prior)):
- self.weights_prior[i] *= self.prior_factor
- return
- @abc.abstractmethod
- def step(self, matrix_pc = None, matrix_neural = None, matrix_npc = None, labels_class = None):
- pass
- class CCCPPCOptimizer(PCOptimizer):
- def __init__(self, pc_joint, pc_marginal, device = torch.device("cuda"), learning_rate = 1e-1, prior_factor = 1e2, projection_epsilon = 1e-2):
- super().__init__(pc_joint, pc_marginal, device, learning_rate, prior_factor, projection_epsilon)
- return
- def step(self, matrix_pc = None, matrix_neural = None, matrix_npc = None, labels_class = None):
- self.pc_joint.reuse_forward = False
- self.pc_marginal.reuse_forward = False
- for sum_node in self.pc_joint.sum_nodes:
- child_values_forward = []
- weight_normalization_sum_node = 0
- for child in sum_node.children:
- child_values_forward.append(child.value_forward)
- child_values_forward = torch.stack(child_values_forward) # number of children x batch size
- weight_updates = torch.exp(sum_node.value_backward + child_values_forward - self.pc_joint.root_node.value_forward) # number of children x batch size
- weight_updates = torch.sum(weight_updates, 1) # number of children
- weight_updates = torch.unsqueeze(weight_updates, 1) # number of children x 1
- sum_node.weights *= weight_updates # number of children x 1
- # Local weight normalization with Laplace smoothing
- for weight_sum_node in sum_node.weights:
- weight_normalization_sum_node += weight_sum_node + self.smoothing_epsilon
- sum_node.weights = (sum_node.weights + self.smoothing_epsilon) / weight_normalization_sum_node
- self.pc_marginal.set_weights(self.pc_joint.get_weights())
- return
- class PGDPCOptimizer(PCOptimizer):
- def __init__(self, pc_joint, pc_marginal, device = torch.device("cuda"), learning_rate = 1e-1, prior_factor = 1e2, projection_epsilon = 1e-2):
- super().__init__(pc_joint, pc_marginal, device, learning_rate, prior_factor, projection_epsilon)
- return
- def step(self, matrix_pc = None, matrix_neural = None, matrix_npc = None, labels_class = None):
- self.pc_joint.reuse_forward = False
- self.pc_marginal.reuse_forward = False
- counter_sum_node = 0
- matrix_pc = matrix_pc.to(self.device)
- matrix_neural = matrix_neural.to(self.device)
- matrix_npc_transposed = matrix_npc.t() # number of class labels x batch size
- matrix_npc_transposed = matrix_npc_transposed.to(self.device)
- labels_class = labels_class.to(self.device)
- progress_bar = None
- if self.device == torch.device("cpu"):
- progress_bar = tqdm.tqdm(total = len(self.pc_joint.sum_nodes), leave = False)
- progress_bar.set_description_str("[INFO]: Optimizing PC")
- for (sum_node_joint, sum_node_marginal, weights_prior) in zip(self.pc_joint.sum_nodes, self.pc_marginal.sum_nodes, self.weights_prior):
- if progress_bar is not None:
- progress_bar.n = counter_sum_node + 1
- progress_bar.refresh()
- counter_sum_node += 1
- for (i, (child_joint, child_marginal)) in enumerate(zip(sum_node_joint.children, sum_node_marginal.children)):
- # Compute weight updates in log space
- weight_updates_joint = torch.exp(sum_node_joint.value_backward + child_joint.value_forward - self.pc_joint.root_node.value_forward)
- weight_updates_marginal = torch.exp(sum_node_marginal.value_backward + child_marginal.value_forward - self.pc_marginal.root_node.value_forward)
- # Set gradients corresponding to zero root node forward values to zero
- mask_1_joint = ((sum_node_joint.value_backward + child_joint.value_forward) == -float("inf"))
- mask_2_joint = (self.pc_joint.root_node.value_forward == -float("inf"))
- mask_joint = mask_1_joint & mask_2_joint
- mask_1_marginal = ((sum_node_marginal.value_backward + child_marginal.value_forward) == -float("inf"))
- mask_2_marginal = (self.pc_marginal.root_node.value_forward == -float("inf"))
- mask_marginal = mask_1_marginal & mask_2_marginal
- weight_updates_joint[mask_joint] = 0
- weight_updates_marginal[mask_marginal] = 0
- weight_updates = weight_updates_joint - weight_updates_marginal
- weight_updates = weight_updates.reshape(matrix_pc.shape) # number of class labels x product of category size of all attributes
- weight_updates *= matrix_pc # number of class labels x product of category size of all attributes
- weight_updates = weight_updates.t() # product of category size of all attributes x number of class labels
- weight_updates = torch.index_select(weight_updates, 1, labels_class) # product of category size of all attributes x batch size
- weight_updates *= matrix_neural # product of category size of all attributes x batch size
- weight_updates = torch.sum(weight_updates, 0) # 1 x batch size
- weight_updates /= matrix_npc_transposed[labels_class, torch.arange(matrix_npc_transposed.shape[1])] # 1 x batch size
- # Average weight updates with Dirichlet prior
- weight_updates_count = weight_updates.shape[0]
- weight_updates = torch.sum(weight_updates, 0, keepdim = True)
- weight_updates += (weights_prior[i] - 1) / sum_node_joint.weights[i]
- weight_updates /= weight_updates_count
- sum_node_joint.weights[i] += self.learning_rate * weight_updates
- if sum_node_joint.weights[i] <= 0:
- sum_node_joint.weights[i] = self.projection_epsilon
- if progress_bar is not None:
- progress_bar.close()
- self.pc_marginal.set_weights(self.pc_joint.get_weights())
- return
- class Node:
- def __init__(self):
- self.children = []
- self.depth = None
- self.device = None
- self.id = None
- self.parents = []
- self.value_backward = None
- self.value_forward = None
- return
- def backward(self):
- if len(self.parents) == 0:
- logger.log_fatal("Node " + str(self.id) + " has no parents. Quit.")
- exit(-1)
- value_backward_parents_product = []
- value_backward_parents_sum = []
- value_backward_product = None
- value_backward_sum = None
- value_forward_parents_product = []
- weights_parents_sum = []
- for parent in self.parents:
- if isinstance(parent, ProductNode):
- value_backward_parents_product.append(parent.value_backward)
- value_forward_parents_product.append(parent.value_forward)
- elif isinstance(parent, SumNode):
- value_backward_parents_sum.append(parent.value_backward)
- weights_parents_sum.append(parent.weights[parent.weights_index_by_child_id[self.id]])
- if len(value_backward_parents_product) > 0:
- value_backward_parents_product = torch.stack(value_backward_parents_product)
- if len(value_backward_parents_sum) > 0:
- value_backward_parents_sum = torch.stack(value_backward_parents_sum)
- if len(value_forward_parents_product) > 0:
- value_forward_parents_product = torch.stack(value_forward_parents_product)
- if len(weights_parents_sum) > 0:
- weights_parents_sum = torch.stack(weights_parents_sum)
- weights_parents_sum = torch.Tensor(weights_parents_sum).reshape(-1, 1)
- weights_parents_sum = weights_parents_sum.to(self.device)
- if len(value_backward_parents_product) > 0:
- value_backward_product = value_backward_parents_product + value_forward_parents_product - self.value_forward
- if len(value_backward_parents_sum) > 0:
- value_backward_sum = value_backward_parents_sum + torch.log(weights_parents_sum)
- if value_backward_product is not None and value_backward_sum is None:
- self.value_backward = value_backward_product
- elif value_backward_product is None and value_backward_sum is not None:
- self.value_backward = value_backward_sum
- else:
- self.value_backward = torch.stack([value_backward_product, value_backward_sum])
- # Compute backward values in log space
- # Log-Sum-Exp trick: https://gregorygundersen.com/blog/2020/02/09/log-sum-exp/
- value_backward_max = torch.max(self.value_backward, 0)[0]
- self.value_backward -= value_backward_max
- self.value_backward = torch.exp(self.value_backward)
- self.value_backward = torch.sum(self.value_backward, 0)
- self.value_backward = torch.log(self.value_backward) + value_backward_max
- return
- @abc.abstractmethod
- def forward(self):
- pass
- class CategoricalLeafNode(Node):
- def __init__(self):
- super().__init__()
- self.attribute_index = None
- self.category_index = None
- return
- def forward(self):
- return
- def set_binary(self, settings):
- if self.attribute_index < 0 or self.category_index < 0:
- logger.log_fatal("Invalid categorical leaf node. Quit.")
- exit(-1)
- categories = settings[self.attribute_index]
- self.value_forward = categories[:, self.category_index].float()
- self.value_forward = self.value_forward.to(self.device)
- # Compute forward values in log space
- self.value_forward = torch.log(self.value_forward)
- return
- def set_categorical(self, settings):
- if self.attribute_index < 0 or self.category_index < 0:
- logger.log_fatal("Invalid categorical leaf node. Quit.")
- exit(-1)
- variables = settings[:, self.attribute_index]
- settings = (variables == self.category_index)
- settings_marginal = (variables < 0)
- self.value_forward = torch.logical_or(settings, settings_marginal).float()
- self.value_forward = self.value_forward.to(self.device)
- # Compute forward values in log space
- self.value_forward = torch.log(self.value_forward)
- return
- class ProductNode(Node):
- def __init__(self):
- super().__init__()
- return
- def forward(self):
- if len(self.children) == 0:
- logger.log_fatal("Product node " + str(self.id) + " has no children. Quit.")
- exit(-1)
- value_forward_children = []
- for child in self.children:
- value_forward_children.append(child.value_forward)
- # Compute forward values in log space
- value_forward_children = torch.stack(value_forward_children)
- self.value_forward = torch.sum(value_forward_children, 0)
- return
- class SumNode(Node):
- def __init__(self):
- super().__init__()
- self.leaf = False
- self.weights = []
- self.weights_index_by_child_id = {}
- return
- def forward(self):
- if len(self.children) == 0:
- logger.log_fatal("Sum node " + str(self.id) + " has no children. Quit.")
- exit(-1)
- value_forward_children = []
- for child in self.children:
- value_forward_children.append(child.value_forward)
- # Compute forward values in log space
- # Log-Sum-Exp trick: https://gregorygundersen.com/blog/2020/02/09/log-sum-exp/
- value_forward_children = torch.stack(value_forward_children)
- value_forward_children_max = torch.max(value_forward_children, 0)[0]
- value_forward_children_max[value_forward_children_max == -float("inf")] = 0
- value_forward_children -= value_forward_children_max
- value_forward_children = torch.exp(value_forward_children)
- value_forward_children *= self.weights
- self.value_forward = torch.sum(value_forward_children, 0)
- self.value_forward = torch.log(self.value_forward) + value_forward_children_max
- return
- class ProbabilisticCircuit:
- def __init__(self, device = torch.device("cuda")):
- self.batch_size = -1
- self.depth = None
- self.device = device
- self.induced_trees = []
- self.leaf_nodes = []
- self.leaf_nodes_dict = {}
- self.nodes = []
- self.product_nodes = []
- self.reuse_backward = False
- self.reuse_forward = False
- self.root_node = None
- self.sum_nodes = []
- self.traversal_order_backward = []
- self.traversal_order_forward = []
- return
- def __call__(self, settings, categorical = True):
- if categorical:
- self.set_leaf_nodes_categorical(settings)
- else:
- self.set_leaf_nodes_binary(settings)
- return self.forward()
- def backward(self):
- if not self.reuse_backward:
- if len(self.traversal_order_backward) == 0:
- logger.log_fatal("Empty tree. Quit.")
- exit(-1)
- if self.root_node is None or self.traversal_order_backward[0][0].id != self.root_node.id:
- logger.log_fatal("Missing root node. Quit.")
- exit(-1)
- if self.root_node.value_forward is None:
- logger.log_fatal("Missing root node forward value. Quit.")
- exit(-1)
- self.reuse_backward = True
- # Initialize root node backward value in log space
- self.root_node.value_backward = torch.log(torch.ones(self.batch_size))
- self.root_node.value_backward = self.root_node.value_backward.to(self.device)
- if self.device == torch.device("cpu"):
- counter_node = 0
- progress_bar = tqdm.tqdm(total = len(self.nodes), leave = False)
- progress_bar.set_description_str("[INFO]: Running PC backward pass")
- for level in self.traversal_order_backward[1:]:
- for node in level:
- progress_bar.n = counter_node + 1
- progress_bar.refresh()
- counter_node += 1
- node.backward()
- progress_bar.close()
- else:
- for level in self.traversal_order_backward[1:]:
- for node in level:
- node.backward()
- return
- def forward(self):
- if not self.reuse_forward:
- if len(self.traversal_order_forward) == 0:
- logger.log_fatal("Empty tree. Quit.")
- exit(-1)
- if self.root_node is None or self.traversal_order_forward[-1][0].id != self.root_node.id:
- logger.log_fatal("Missing root node. Quit.")
- exit(-1)
- self.reuse_backward = False
- self.reuse_forward = True
- if self.device == torch.device("cpu"):
- counter_node = 0
- progress_bar = tqdm.tqdm(total = len(self.nodes), leave = False)
- progress_bar.set_description_str("[INFO]: Running PC forward pass")
- for level in self.traversal_order_forward:
- for node in level:
- progress_bar.n = counter_node + 1
- progress_bar.refresh()
- counter_node += 1
- node.forward()
- progress_bar.close()
- else:
- for level in self.traversal_order_forward:
- for node in level:
- node.forward()
- return self.root_node.value_forward
- def gather_induced_trees(self):
- for induced_tree in self.recurse_induced_trees(self.root_node):
- induced_tree[1] = numpy.prod(induced_tree[1])
- self.induced_trees.append(induced_tree)
- return
- def get_weights(self):
- weights = []
- for sum_node in self.sum_nodes:
- weights.append(torch.clone(sum_node.weights))
- return weights
- def load(self, file_path_pc):
- if not os.path.exists(file_path_pc):
- logger.log_fatal("Invalid PC file path. Quit.")
- exit(-1)
- self.reuse_backward = False
- self.reuse_forward = False
- with open(file_path_pc, "r") as file_pc:
- categorical_leaf_node_id = -1
- counter_line = 0
- id_to_nodes = {}
- lines = file_pc.readlines()
- progress_bar = tqdm.tqdm(total = len(lines), leave = False)
- reading_nodes = True
- progress_bar.set_description_str("[INFO]: Loading PC")
- for line in lines:
- progress_bar.n = counter_line + 1
- progress_bar.refresh()
- counter_line += 1
- line = line.strip()
- if line[0] == "#":
- line = line.replace("#", "")
- if line == "NODES":
- reading_nodes = True
- elif line == "EDGES":
- reading_nodes = False
- continue
- line_list = line.split(",")
- if reading_nodes:
- node_id = int(line_list[0])
- node_type = line_list[1]
- if node_type == "SUM":
- sum_node = SumNode()
- sum_node.device = self.device
- sum_node.id = node_id
- self.nodes.append(sum_node)
- id_to_nodes[sum_node.id] = sum_node
- self.sum_nodes.append(sum_node)
- elif node_type == "PRD":
- product_node = ProductNode()
- product_node.device = self.device
- product_node.id = node_id
- self.nodes.append(product_node)
- id_to_nodes[product_node.id] = product_node
- self.product_nodes.append(product_node)
- elif node_type == "CatNode" or node_type == "CATNODE":
- categorical_leaf_node_list = []
- node_attribute_index = int(line_list[2])
- node_probabilities = line_list[3:]
- for i in range(0, len(node_probabilities)):
- node_probabilities[i] = float(node_probabilities[i])
- if node_attribute_index in self.leaf_nodes_dict.keys():
- categorical_leaf_node_list = self.leaf_nodes_dict[node_attribute_index]
- else:
- for category_index in range(0, len(node_probabilities)):
- categorical_leaf_node = CategoricalLeafNode()
- categorical_leaf_node.attribute_index = node_attribute_index
- categorical_leaf_node.category_index = category_index
- categorical_leaf_node.device = self.device
- categorical_leaf_node.id = categorical_leaf_node_id
- categorical_leaf_node_id -= 1
- categorical_leaf_node_list.append(categorical_leaf_node)
- self.leaf_nodes.append(categorical_leaf_node)
- self.nodes.append(categorical_leaf_node)
- id_to_nodes[categorical_leaf_node.id] = categorical_leaf_node
- self.leaf_nodes_dict[node_attribute_index] = categorical_leaf_node_list
- sum_node = SumNode()
- sum_node.children = categorical_leaf_node_list
- sum_node.device = self.device
- sum_node.id = node_id
- sum_node.leaf = True
- sum_node.weights = node_probabilities
- self.nodes.append(sum_node)
- id_to_nodes[sum_node.id] = sum_node
- self.sum_nodes.append(sum_node)
- for (i, categorical_leaf_node) in enumerate(categorical_leaf_node_list):
- sum_node.weights_index_by_child_id[categorical_leaf_node.id] = i
- categorical_leaf_node.parents.append(sum_node)
- elif node_type == "CATNODEPRD":
- node_attribute_index = int(line_list[2])
- node_category_index = int(line_list[3])
- categorical_leaf_node = CategoricalLeafNode()
- categorical_leaf_node.attribute_index = node_attribute_index
- categorical_leaf_node.category_index = node_category_index
- categorical_leaf_node.device = self.device
- categorical_leaf_node.id = node_id
- self.leaf_nodes.append(categorical_leaf_node)
- self.nodes.append(categorical_leaf_node)
- id_to_nodes[categorical_leaf_node.id] = categorical_leaf_node
- else:
- nodes = []
- node_id_first = int(line_list[0])
- node_id_second = int(line_list[1])
- nodes.append(id_to_nodes[node_id_first])
- nodes.append(id_to_nodes[node_id_second])
- if len(nodes) != 2:
- logger.log_fatal("Invalid edge. Quit.")
- exit(-1)
- if len(line_list) >= 3:
- node_weight = float(line_list[2])
- if isinstance(nodes[0], SumNode) and not nodes[0].leaf:
- nodes[0].children.append(nodes[1])
- nodes[0].weights.append(node_weight)
- nodes[0].weights_index_by_child_id[nodes[1].id] = len(nodes[0].weights) - 1
- nodes[1].parents.append(nodes[0])
- elif isinstance(nodes[1], SumNode) and not nodes[0].leaf:
- nodes[1].children.append(nodes[0])
- nodes[1].weights.append(node_weight)
- nodes[1].weights_index_by_child_id[nodes[0].id] = len(nodes[1].weights) - 1
- nodes[0].parents.append(nodes[1])
- else:
- if isinstance(nodes[0], ProductNode):
- nodes[0].children.append(nodes[1])
- nodes[1].parents.append(nodes[0])
- elif isinstance(nodes[1], ProductNode):
- nodes[1].children.append(nodes[0])
- nodes[0].parents.append(nodes[1])
- progress_bar.close()
- for sum_node in self.sum_nodes:
- sum_node.weights = torch.Tensor(sum_node.weights).reshape(-1, 1)
- sum_node.weights = sum_node.weights.to(self.device)
- root_nodes = []
- for node in self.nodes:
- if len(node.parents) == 0:
- root_nodes.append(node)
- if len(root_nodes) != 1:
- logger.log_fatal("Invalid PC. Quit.")
- exit(-1)
- self.depth = self.traverse({0: root_nodes}, 0)
- self.root_node = root_nodes[0]
- self.traversal_order_backward = self.topological_sort_backward()
- self.traversal_order_forward = self.traversal_order_backward.copy()
- self.traversal_order_forward.reverse()
- def getLengthNestedList(nested_list):
- length = 0
- for item in nested_list:
- if isinstance(item, list):
- length += getLengthNestedList(item)
- else:
- length += 1
- return length
- if getLengthNestedList(self.traversal_order_backward) != len(self.nodes):
- logger.log_fatal("Invalid PC backward traversal. Quit.")
- exit(-1)
- if getLengthNestedList(self.traversal_order_forward) != len(self.nodes):
- logger.log_fatal("Invalid PC forward traversal. Quit.")
- exit(-1)
- return
- def normalize_weights(self, smoothing_epsilon):
- for sum_node in self.sum_nodes:
- value_forward_children = []
- weight_normalization = 0
- for child in sum_node.children:
- value_forward_children.append(child.value_forward)
- value_forward_children = torch.stack(value_forward_children)
- value_forward_children_max = torch.max(value_forward_children, 0)[0]
- for (i, child) in enumerate(sum_node.children):
- weight_normalization += sum_node.weights[i] * torch.exp(child.value_forward - value_forward_children_max) + smoothing_epsilon
- # Local weight normalization with Laplace smoothing
- for (i, child) in enumerate(sum_node.children):
- weight = sum_node.weights[i] * torch.exp(child.value_forward - value_forward_children_max) + smoothing_epsilon
- sum_node.weights[i] = weight / weight_normalization
- return
- def randomize_weights(self):
- self.reuse_backward = False
- self.reuse_forward = False
- weights = self.get_weights()
- for i in range(len(weights)):
- weights[i] = weights[i].uniform_(0, 1)
- weights[i] /= torch.sum(weights[i])
- self.set_weights(weights)
- return
- def recurse_induced_trees(self, node):
- if isinstance(node, SumNode):
- for (i, child) in enumerate(node.children):
- for sub_induced_tree in self.recurse_induced_trees(child):
- induced_tree = [{node}, [node.weights[i].item()]]
- induced_tree[0].update(sub_induced_tree[0])
- induced_tree[1] += sub_induced_tree[1]
- yield induced_tree
- elif isinstance(node, ProductNode):
- sub_induced_trees_product = []
- for child in node.children:
- sub_induced_trees = []
- for sub_induced_tree in self.recurse_induced_trees(child):
- sub_induced_trees.append(sub_induced_tree)
- sub_induced_trees_product.append(sub_induced_trees)
- for sub_induced_trees in itertools.product(*sub_induced_trees_product):
- induced_tree = [{node}, []]
- for sub_induced_tree in sub_induced_trees:
- induced_tree[0].update(sub_induced_tree[0])
- induced_tree[1] += sub_induced_tree[1]
- yield induced_tree
- elif isinstance(node, CategoricalLeafNode):
- yield [{node}, []]
- return
- def set_leaf_nodes_binary(self, settings):
- self.reuse_backward = False
- self.reuse_forward = False
- self.batch_size = settings[0].shape[0]
- for leaf_node in self.leaf_nodes:
- leaf_node.set_binary(settings)
- return
- def set_leaf_nodes_categorical(self, settings):
- self.reuse_backward = False
- self.reuse_forward = False
- self.batch_size = settings.shape[0]
- for leaf_node in self.leaf_nodes:
- leaf_node.set_categorical(settings)
- return
- def set_weights(self, weights):
- if len(weights) != len(self.sum_nodes):
- logger.log_fatal("Invalid weights. Quit.")
- exit(-1)
- self.reuse_backward = False
- self.reuse_forward = False
- for (sum_node, sum_node_weights) in zip(self.sum_nodes, weights):
- sum_node.weights = torch.clone(sum_node_weights).to(self.device)
- return
- def topological_sort_backward(self):
- levels = []
- parent_count = {}
- queue = []
- for node in self.nodes:
- parent_count[node.id] = len(node.parents)
- queue.append(self.root_node)
- while len(queue) > 0:
- level = []
- queue_next_level = []
- for node in queue:
- level.append(node)
- for child in node.children:
- parent_count[child.id] -= 1
- if parent_count[child.id] <= 0:
- queue_next_level.append(child)
- levels.append(level)
- queue = queue_next_level.copy()
- return levels
- def traverse(self, layers, depth):
- if depth not in layers.keys():
- return depth - 1
- for node in layers[depth]:
- node.depth = depth
- for child in node.children:
- if depth + 1 not in layers.keys():
- layers[depth + 1] = []
- layers[depth + 1].append(child)
- depth_next = self.traverse(layers, depth + 1)
- if depth_next > depth:
- depth = depth_next
- return depth
pc.py at commit 4a46c28, under CC-BY-NC-SA-4.0 · at the source
Overview
- Department of Computer Science, University of Illinois Urbana-Champaign, Urbana, IL USA
- Department of Electrical and Computer Engineering, University of Illinois Urbana-Champaign, Urbana, IL USA
- Department of Computer Science, College of William and Mary, Williamsburg, VA USA
Abstract
End-to-end deep neural networks have achieved remarkable success across various domains but are often criticized for their lack of interpretability. While post hoc explanation methods attempt to address this issue, they often fail to accurately represent these black-box models, resulting in misleading or incomplete explanations. To overcome these challenges, we propose an inherently transparent model architecture called Neural Probabilistic Circuits (NPCs), which enable compositional and interpretable predictions through logical reasoning. In particular, an NPC consists of two modules: an attribute recognition model, which predicts probabilities for various attributes, and a task predictor built on a probabilistic circuit, which enables logical reasoning over recognized attributes to make class predictions. To train NPCs, we introduce a three-stage training algorithm comprising attribute recognition, circuit construction, and joint optimization. Moreover, we theoretically demonstrate that an NPC’s error is upper-bounded by a linear combination of the errors from its modules. To further demonstrate the interpretability of NPC, we provide both the most probable explanations and the counterfactual explanations. Empirical results on four benchmark datasets show that NPCs strike a balance between interpretability and performance, achieving results competitive even with those of end-to-end black-box models while providing enhanced interpretability.
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 8 matches between paragraphs and lines of code.
uiuctml/npc-models
4a46c28f7cbfcc4fd89c5460de74c573fd14002c, 9 October 2025Availability: 1 check, the latest on 27 September 2026: the link answers
- 27 September 2026: the link answers
30 files
- scripts/
npc-models/ , Shell, 95 linestest_baseline.bash - scripts/
npc-models/ , Shell, 35 linestest_blackbox.bash - scripts/
npc-models/ , Shell, 80 linestest_neural.bash - scripts/
npc-models/ , Shell, 80 linestest_npc.bash - scripts/
npc-models/ , Shell, 80 linestest_pc.bash - scripts/
npc-models/ , Shell, 35 linestrain_baseline.bash - scripts/
npc-models/ , Shell, 20 linestrain_blackbox.bash - scripts/
npc-models/ , Shell, 20 linestrain_neural.bash - scripts/
npc-models/ , Shell, 58 lines, 1 matchtrain_npc.bash - scripts/
npc-models/ , Shell, 20 linestrain_pc.bash - src/
npc-models/ , Python, 135 linesdataset.py - src/
npc-models/ , Python, 150 lines, 1 matchheader.py - src/
npc-models/ , Python, 438 linesinterpret.py - src/
npc-models/ , Python, 87 lineslogger.py - src/
npc-models/ , Python, 207 linesmodel.py - src/
npc-models/ , Python, 829 lines, 2 matchespc.py - src/
npc-models/ , Python, 210 linestest_baseline.py - src/
npc-models/ , Python, 120 linestest_blackbox.py - src/
npc-models/ , Python, 139 linestest_neural.py - src/
npc-models/ , Python, 519 linestest_npc.py - src/
npc-models/ , Python, 122 linestest_pc.py - src/
npc-models/ , Python, 264 linestrain_baseline.py - src/
npc-models/ , Python, 211 linestrain_blackbox.py - src/
npc-models/ , Python, 222 linestrain_neural.py - src/
npc-models/ , Python, 340 lines, 2 matchestrain_npc.py - src/
npc-models/ , Python, 154 lines, 2 matchestrain_pc.py - src/
npc-models/ , Python, 44 linestype.py - src/
npc-models/ , Python, 190 linesutility.py - LICENSE, License, 437 lines
- README.md, Text, 337 lines
Code Availability
The code is available at https://
Reproduced under the paper's license (CC BY), from the paper cited above.
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;
- 28 scripts, each with its path and the digest of its content;
- 8 matches 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
- huggingface.co/
datasets/ , at Hugging Face; found in “Data Availability”ylecun/ mnist - kaggle.com/
datasets/ , at Kaggle; found in “Data Availability”meowmeowmeowmeowmeow
Data Availability
We employed four publicly available image datasets in this work, namely, Animal with Attributes 2 (AwA2), CelebFaces Attributes (CelebA), German Traffic Sign Recognition Benchmark (GTSRB), and Modified National Institute of Standards and Technology (MNIST). AwA2 is available at https://
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 2, 28 September 2026
- Publisher: n/a → Springer Science+Business Media
- Funding: added National Science Foundation: 2311085, #80NSSC22M0070; National Aeronautics and Space Administration: 80NSSC22M0070; Division of Electrical, Communications and Cyber Systems
Version 1, 27 September 2026: the first record
Recorded: type, language, journal, volume, issue, pages, dates, 5 authors, 4 keywords, 12 references.
Cite
This paper
Chen, W., Yu, S., Shao, H., Sha, L., & Zhao, H. (2026). Neural Probabilistic Circuits: Enabling Compositional and Interpretable Predictions Through Logical Reasoning. Machine learning, 115(9), 207. https://
BibTeX
@article{chen2026neural,
author = {Chen, Weixin and Yu, Simon and Shao, Huajie and Sha, Lui and Zhao, Han},
title = {{Neural Probabilistic Circuits: Enabling Compositional and Interpretable Predictions Through Logical Reasoning}},
journal = {Machine learning},
year = {2026},
month = aug,
volume = {115},
number = {9},
pages = {207},
publisher = {Springer Science+Business Media},
issn = {0885-6125},
doi = {10.1007/
url = {https://
pmid = {42670317},
pmcid = {PMC13525958}
}
RIS
TY - JOUR
AU - Chen, Weixin
AU - Yu, Simon
AU - Shao, Huajie
AU - Sha, Lui
AU - Zhao, Han
TI - Neural Probabilistic Circuits: Enabling Compositional and Interpretable Predictions Through Logical Reasoning
T2 - Machine learning
J2 - Mach Learn
PY - 2026
DA - 2026/
VL - 115
IS - 9
SP - 207
SN - 0885-6125
PB - Springer Science+Business Media
DO - 10.1007/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.1007/
"type": "article-journal",
"title": "Neural Probabilistic Circuits: Enabling Compositional and Interpretable Predictions Through Logical Reasoning",
"container-title": "Machine learning",
"author": [
{
"family": "Chen",
"given": "Weixin"
},
{
"family": "Yu",
"given": "Simon"
},
{
"family": "Shao",
"given": "Huajie"
},
{
"family": "Sha",
"given": "Lui"
},
{
"family": "Zhao",
"given": "Han"
}
],
"container-title-short":
"volume": "115",
"issue": "9",
"page": "207",
"DOI": "10.1007/
"PMID": "42670317",
"PMCID": "PMC13525958",
"ISSN": "0885-6125",
"publisher": "Springer Science+Business Media",
"URL": "https://
"language": "en",
"issued": {
"date-parts": [
[
2026,
8,
29
]
]
}
}
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.1093/braincomms/fcag253 [code]
- Disease detection and classification in temporal lobe epilepsy: step-wise versus simultaneous AI decision models in a multisite neuroimaging study.Journal: Brain communicationsIn common: Pillow, PyTorch, scikit-learn, 1 other tool, 1 reference
- [2] doi:10.1177/26331055261460858 [code]
- The Geometric Signatures of Brain State Transitions: Recursive Informational Curvature Reveals Hidden Dynamics in Primate Cortex.Journal: Neuroscience insightsIn common: Pillow, scikit-learn, NumPy, systems, 1 reference
- [3] doi:10.1016/j.isci.2026.117068 [code]
- Directed graph neural networks with partial directed coherence for seizure prediction and epileptogenic network characterization.Journal: iScienceIn common: PyTorch, scikit-learn, NumPy, systems, 1 reference
- [4] doi:10.1371/journal.pcbi.1014656 [code]
- Contrastive learning to fine-tune feature extraction models for the visual cortex.Journal: PLoS computational biologyIn common: Pillow, PyTorch, scikit-learn, 1 other tool, systems
- [5] doi:10.1038/s41593-026-02388-9 [code]
- Hippocampal CA3 connectomics reveals a gradient of mossy fiber inputs and selective feedforward inhibition onto pyramidal cells.Journal: Nature neuroscienceIn common: Pillow, PyTorch, scikit-learn, 1 other tool, systems
- [6] doi:10.1523/jneurosci.0038-26.2026 [code]
- Multidimensional Feature Tuning in Category Selective Areas of Human Visual Cortex.Journal: The Journal of neuroscience : the official journal of the Society for NeuroscienceIn common: Pillow, PyTorch, scikit-learn, 1 other tool, systems
- [7] doi:10.1016/j.isci.2026.116825 [code]
- Social hierarchy shapes behavioral and transcriptional responses to chronic stress and ketamine in male mice.Journal: iScienceIn common: Pillow, PyTorch, scikit-learn, 1 other tool, systems
- [8] doi:10.1038/s41467-026-75347-4 [code]
- Sleep reveals dynamics integrating and segregating movement and stimulus representations in V1.Journal: Nature communicationsIn common: Pillow, PyTorch, scikit-learn, 1 other tool, systems
- [9] doi:10.1016/j.isci.2026.116206 [code]
- Gut distension evokes rapid neural dynamics in vagal and hindbrain populations of larval zebrafish.Journal: iScienceIn common: Pillow, PyTorch, scikit-learn, 1 other tool, systems
- [10] doi:10.1371/journal.pbio.3003824 [code]
- Flexible goal learning involves coordinated population activity in dCA1 and medial orbitofrontal cortex.Journal: PLoS biologyIn common: Pillow, PyTorch, scikit-learn, 1 other tool, systems
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, 28 scripts, and 8 matches 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:5fa15d84b2c719ea…
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.
