OSCR

Neural Probabilistic Circuits: Enabling Compositional and Interpretable Predictions Through Logical Reasoning.

Code ↔ Paper

8 matches between paragraphs of the paper and lines of its authors' code, computed by the harvester (lexical-v1). Click a colored paragraph or line to see its counterpart.

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. [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. [2] § Preliminaries ↔ src/npc-models/pc.py, lines 113–144 · score 0.67 · weighted sum, weight updating, root node, sum node, children, joint
  3. [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. [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. [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. [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. [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. [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

  1. """
  2. @file pc.py
  3. @author Simon Yu
  4. @date 02/07/2024
  5. @brief PC classes.
  6. """
  7. import abc
  8. import itertools
  9. import logger
  10. import numpy
  11. import os
  12. import torch
  13. import tqdm
  14. class PCLearningRateScheduler:
  15. def __init__(self, optimizer, factor = 0.8, patience = 2, threshold = 1e-4, cooldown = 2, min_learning_rate = 1e-6):
  16. self.cooldown = cooldown
  17. self.cooldown_counter = 0
  18. self.factor = factor
  19. self.metrics = None
  20. self.min_learning_rate = min_learning_rate
  21. self.optimizer = optimizer
  22. self.patience = patience
  23. self.patience_counter = 0
  24. self.threshold = threshold
  25. return
  26. @abc.abstractmethod
  27. def step(self, metrics):
  28. pass
  29. class LikelihoodPCLearningRateScheduler(PCLearningRateScheduler):
  30. def __init__(self, optimizer, factor = 0.8):
  31. super().__init__(optimizer, factor)
  32. return
  33. def step(self, metrics):
  34. if self.metrics is not None:
  35. if metrics < self.metrics:
  36. self.optimizer.learning_rate *= self.factor
  37. logger.log_info("Reducing PC learning rate to " + "{:e}".format(self.optimizer.learning_rate) + "...")
  38. self.metrics = metrics
  39. return
  40. class LossPCLearningRateScheduler(PCLearningRateScheduler):
  41. def __init__(self, optimizer, factor = 0.8, patience = 2, threshold = 1e-4, cooldown = 2, min_learning_rate = 1e-6):
  42. super().__init__(optimizer, factor, patience, threshold, cooldown, min_learning_rate)
  43. return
  44. def step(self, metrics):
  45. if self.metrics is None:
  46. self.metrics = metrics
  47. return
  48. if self.optimizer.learning_rate < self.min_learning_rate:
  49. self.metrics = metrics
  50. return
  51. if self.cooldown_counter > 0:
  52. self.cooldown_counter -= 1
  53. self.metrics = metrics
  54. return
  55. if abs(self.metrics - metrics) < self.threshold:
  56. if self.patience_counter < self.patience:
  57. self.patience_counter += 1
  58. self.metrics = metrics
  59. return
  60. self.optimizer.learning_rate *= self.factor
  61. logger.log_info("Reducing PC learning rate to " + "{:e}".format(self.optimizer.learning_rate) + "...")
  62. self.cooldown_counter = self.cooldown
  63. self.patience_counter = 0
  64. else:
  65. self.patience_counter = 0
  66. self.metrics = metrics
  67. return
  68. class PCOptimizer:
  69. def __init__(self, pc_joint, pc_marginal, device = torch.device("cuda"), learning_rate = 1e-1, prior_factor = 1e2, projection_epsilon = 1e-2):
  70. self.device = device
  71. self.learning_rate = learning_rate
  72. self.smoothing_epsilon = torch.finfo(torch.float).eps
  73. self.prior_factor = prior_factor
  74. self.projection_epsilon = projection_epsilon
  75. self.pc_joint = pc_joint
  76. self.pc_marginal = pc_marginal
  77. self.weights_prior = None
  78. return
  79. def set_weights_prior(self, weights_prior):
  80. self.weights_prior = weights_prior
  81. for i in range(len(self.weights_prior)):
  82. self.weights_prior[i] *= self.prior_factor
  83. return
  84. @abc.abstractmethod
  85. def step(self, matrix_pc = None, matrix_neural = None, matrix_npc = None, labels_class = None):
  86. pass
  87. class CCCPPCOptimizer(PCOptimizer):
  88. def __init__(self, pc_joint, pc_marginal, device = torch.device("cuda"), learning_rate = 1e-1, prior_factor = 1e2, projection_epsilon = 1e-2):
  89. super().__init__(pc_joint, pc_marginal, device, learning_rate, prior_factor, projection_epsilon)
  90. return
  91. def step(self, matrix_pc = None, matrix_neural = None, matrix_npc = None, labels_class = None):
  92. self.pc_joint.reuse_forward = False
  93. self.pc_marginal.reuse_forward = False
  94. for sum_node in self.pc_joint.sum_nodes:
  95. child_values_forward = []
  96. weight_normalization_sum_node = 0
  97. for child in sum_node.children:
  98. child_values_forward.append(child.value_forward)
  99. child_values_forward = torch.stack(child_values_forward) # number of children x batch size
  100. weight_updates = torch.exp(sum_node.value_backward + child_values_forward - self.pc_joint.root_node.value_forward) # number of children x batch size
  101. weight_updates = torch.sum(weight_updates, 1) # number of children
  102. weight_updates = torch.unsqueeze(weight_updates, 1) # number of children x 1
  103. sum_node.weights *= weight_updates # number of children x 1
  104. # Local weight normalization with Laplace smoothing
  105. for weight_sum_node in sum_node.weights:
  106. weight_normalization_sum_node += weight_sum_node + self.smoothing_epsilon
  107. sum_node.weights = (sum_node.weights + self.smoothing_epsilon) / weight_normalization_sum_node
  108. self.pc_marginal.set_weights(self.pc_joint.get_weights())
  109. return
  110. class PGDPCOptimizer(PCOptimizer):
  111. def __init__(self, pc_joint, pc_marginal, device = torch.device("cuda"), learning_rate = 1e-1, prior_factor = 1e2, projection_epsilon = 1e-2):
  112. super().__init__(pc_joint, pc_marginal, device, learning_rate, prior_factor, projection_epsilon)
  113. return
  114. def step(self, matrix_pc = None, matrix_neural = None, matrix_npc = None, labels_class = None):
  115. self.pc_joint.reuse_forward = False
  116. self.pc_marginal.reuse_forward = False
  117. counter_sum_node = 0
  118. matrix_pc = matrix_pc.to(self.device)
  119. matrix_neural = matrix_neural.to(self.device)
  120. matrix_npc_transposed = matrix_npc.t() # number of class labels x batch size
  121. matrix_npc_transposed = matrix_npc_transposed.to(self.device)
  122. labels_class = labels_class.to(self.device)
  123. progress_bar = None
  124. if self.device == torch.device("cpu"):
  125. progress_bar = tqdm.tqdm(total = len(self.pc_joint.sum_nodes), leave = False)
  126. progress_bar.set_description_str("[INFO]: Optimizing PC")
  127. for (sum_node_joint, sum_node_marginal, weights_prior) in zip(self.pc_joint.sum_nodes, self.pc_marginal.sum_nodes, self.weights_prior):
  128. if progress_bar is not None:
  129. progress_bar.n = counter_sum_node + 1
  130. progress_bar.refresh()
  131. counter_sum_node += 1
  132. for (i, (child_joint, child_marginal)) in enumerate(zip(sum_node_joint.children, sum_node_marginal.children)):
  133. # Compute weight updates in log space
  134. weight_updates_joint = torch.exp(sum_node_joint.value_backward + child_joint.value_forward - self.pc_joint.root_node.value_forward)
  135. weight_updates_marginal = torch.exp(sum_node_marginal.value_backward + child_marginal.value_forward - self.pc_marginal.root_node.value_forward)
  136. # Set gradients corresponding to zero root node forward values to zero
  137. mask_1_joint = ((sum_node_joint.value_backward + child_joint.value_forward) == -float("inf"))
  138. mask_2_joint = (self.pc_joint.root_node.value_forward == -float("inf"))
  139. mask_joint = mask_1_joint & mask_2_joint
  140. mask_1_marginal = ((sum_node_marginal.value_backward + child_marginal.value_forward) == -float("inf"))
  141. mask_2_marginal = (self.pc_marginal.root_node.value_forward == -float("inf"))
  142. mask_marginal = mask_1_marginal & mask_2_marginal
  143. weight_updates_joint[mask_joint] = 0
  144. weight_updates_marginal[mask_marginal] = 0
  145. weight_updates = weight_updates_joint - weight_updates_marginal
  146. weight_updates = weight_updates.reshape(matrix_pc.shape) # number of class labels x product of category size of all attributes
  147. weight_updates *= matrix_pc # number of class labels x product of category size of all attributes
  148. weight_updates = weight_updates.t() # product of category size of all attributes x number of class labels
  149. weight_updates = torch.index_select(weight_updates, 1, labels_class) # product of category size of all attributes x batch size
  150. weight_updates *= matrix_neural # product of category size of all attributes x batch size
  151. weight_updates = torch.sum(weight_updates, 0) # 1 x batch size
  152. weight_updates /= matrix_npc_transposed[labels_class, torch.arange(matrix_npc_transposed.shape[1])] # 1 x batch size
  153. # Average weight updates with Dirichlet prior
  154. weight_updates_count = weight_updates.shape[0]
  155. weight_updates = torch.sum(weight_updates, 0, keepdim = True)
  156. weight_updates += (weights_prior[i] - 1) / sum_node_joint.weights[i]
  157. weight_updates /= weight_updates_count
  158. sum_node_joint.weights[i] += self.learning_rate * weight_updates
  159. if sum_node_joint.weights[i] <= 0:
  160. sum_node_joint.weights[i] = self.projection_epsilon
  161. if progress_bar is not None:
  162. progress_bar.close()
  163. self.pc_marginal.set_weights(self.pc_joint.get_weights())
  164. return
  165. class Node:
  166. def __init__(self):
  167. self.children = []
  168. self.depth = None
  169. self.device = None
  170. self.id = None
  171. self.parents = []
  172. self.value_backward = None
  173. self.value_forward = None
  174. return
  175. def backward(self):
  176. if len(self.parents) == 0:
  177. logger.log_fatal("Node " + str(self.id) + " has no parents. Quit.")
  178. exit(-1)
  179. value_backward_parents_product = []
  180. value_backward_parents_sum = []
  181. value_backward_product = None
  182. value_backward_sum = None
  183. value_forward_parents_product = []
  184. weights_parents_sum = []
  185. for parent in self.parents:
  186. if isinstance(parent, ProductNode):
  187. value_backward_parents_product.append(parent.value_backward)
  188. value_forward_parents_product.append(parent.value_forward)
  189. elif isinstance(parent, SumNode):
  190. value_backward_parents_sum.append(parent.value_backward)
  191. weights_parents_sum.append(parent.weights[parent.weights_index_by_child_id[self.id]])
  192. if len(value_backward_parents_product) > 0:
  193. value_backward_parents_product = torch.stack(value_backward_parents_product)
  194. if len(value_backward_parents_sum) > 0:
  195. value_backward_parents_sum = torch.stack(value_backward_parents_sum)
  196. if len(value_forward_parents_product) > 0:
  197. value_forward_parents_product = torch.stack(value_forward_parents_product)
  198. if len(weights_parents_sum) > 0:
  199. weights_parents_sum = torch.stack(weights_parents_sum)
  200. weights_parents_sum = torch.Tensor(weights_parents_sum).reshape(-1, 1)
  201. weights_parents_sum = weights_parents_sum.to(self.device)
  202. if len(value_backward_parents_product) > 0:
  203. value_backward_product = value_backward_parents_product + value_forward_parents_product - self.value_forward
  204. if len(value_backward_parents_sum) > 0:
  205. value_backward_sum = value_backward_parents_sum + torch.log(weights_parents_sum)
  206. if value_backward_product is not None and value_backward_sum is None:
  207. self.value_backward = value_backward_product
  208. elif value_backward_product is None and value_backward_sum is not None:
  209. self.value_backward = value_backward_sum
  210. else:
  211. self.value_backward = torch.stack([value_backward_product, value_backward_sum])
  212. # Compute backward values in log space
  213. # Log-Sum-Exp trick: https://gregorygundersen.com/blog/2020/02/09/log-sum-exp/
  214. value_backward_max = torch.max(self.value_backward, 0)[0]
  215. self.value_backward -= value_backward_max
  216. self.value_backward = torch.exp(self.value_backward)
  217. self.value_backward = torch.sum(self.value_backward, 0)
  218. self.value_backward = torch.log(self.value_backward) + value_backward_max
  219. return
  220. @abc.abstractmethod
  221. def forward(self):
  222. pass
  223. class CategoricalLeafNode(Node):
  224. def __init__(self):
  225. super().__init__()
  226. self.attribute_index = None
  227. self.category_index = None
  228. return
  229. def forward(self):
  230. return
  231. def set_binary(self, settings):
  232. if self.attribute_index < 0 or self.category_index < 0:
  233. logger.log_fatal("Invalid categorical leaf node. Quit.")
  234. exit(-1)
  235. categories = settings[self.attribute_index]
  236. self.value_forward = categories[:, self.category_index].float()
  237. self.value_forward = self.value_forward.to(self.device)
  238. # Compute forward values in log space
  239. self.value_forward = torch.log(self.value_forward)
  240. return
  241. def set_categorical(self, settings):
  242. if self.attribute_index < 0 or self.category_index < 0:
  243. logger.log_fatal("Invalid categorical leaf node. Quit.")
  244. exit(-1)
  245. variables = settings[:, self.attribute_index]
  246. settings = (variables == self.category_index)
  247. settings_marginal = (variables < 0)
  248. self.value_forward = torch.logical_or(settings, settings_marginal).float()
  249. self.value_forward = self.value_forward.to(self.device)
  250. # Compute forward values in log space
  251. self.value_forward = torch.log(self.value_forward)
  252. return
  253. class ProductNode(Node):
  254. def __init__(self):
  255. super().__init__()
  256. return
  257. def forward(self):
  258. if len(self.children) == 0:
  259. logger.log_fatal("Product node " + str(self.id) + " has no children. Quit.")
  260. exit(-1)
  261. value_forward_children = []
  262. for child in self.children:
  263. value_forward_children.append(child.value_forward)
  264. # Compute forward values in log space
  265. value_forward_children = torch.stack(value_forward_children)
  266. self.value_forward = torch.sum(value_forward_children, 0)
  267. return
  268. class SumNode(Node):
  269. def __init__(self):
  270. super().__init__()
  271. self.leaf = False
  272. self.weights = []
  273. self.weights_index_by_child_id = {}
  274. return
  275. def forward(self):
  276. if len(self.children) == 0:
  277. logger.log_fatal("Sum node " + str(self.id) + " has no children. Quit.")
  278. exit(-1)
  279. value_forward_children = []
  280. for child in self.children:
  281. value_forward_children.append(child.value_forward)
  282. # Compute forward values in log space
  283. # Log-Sum-Exp trick: https://gregorygundersen.com/blog/2020/02/09/log-sum-exp/
  284. value_forward_children = torch.stack(value_forward_children)
  285. value_forward_children_max = torch.max(value_forward_children, 0)[0]
  286. value_forward_children_max[value_forward_children_max == -float("inf")] = 0
  287. value_forward_children -= value_forward_children_max
  288. value_forward_children = torch.exp(value_forward_children)
  289. value_forward_children *= self.weights
  290. self.value_forward = torch.sum(value_forward_children, 0)
  291. self.value_forward = torch.log(self.value_forward) + value_forward_children_max
  292. return
  293. class ProbabilisticCircuit:
  294. def __init__(self, device = torch.device("cuda")):
  295. self.batch_size = -1
  296. self.depth = None
  297. self.device = device
  298. self.induced_trees = []
  299. self.leaf_nodes = []
  300. self.leaf_nodes_dict = {}
  301. self.nodes = []
  302. self.product_nodes = []
  303. self.reuse_backward = False
  304. self.reuse_forward = False
  305. self.root_node = None
  306. self.sum_nodes = []
  307. self.traversal_order_backward = []
  308. self.traversal_order_forward = []
  309. return
  310. def __call__(self, settings, categorical = True):
  311. if categorical:
  312. self.set_leaf_nodes_categorical(settings)
  313. else:
  314. self.set_leaf_nodes_binary(settings)
  315. return self.forward()
  316. def backward(self):
  317. if not self.reuse_backward:
  318. if len(self.traversal_order_backward) == 0:
  319. logger.log_fatal("Empty tree. Quit.")
  320. exit(-1)
  321. if self.root_node is None or self.traversal_order_backward[0][0].id != self.root_node.id:
  322. logger.log_fatal("Missing root node. Quit.")
  323. exit(-1)
  324. if self.root_node.value_forward is None:
  325. logger.log_fatal("Missing root node forward value. Quit.")
  326. exit(-1)
  327. self.reuse_backward = True
  328. # Initialize root node backward value in log space
  329. self.root_node.value_backward = torch.log(torch.ones(self.batch_size))
  330. self.root_node.value_backward = self.root_node.value_backward.to(self.device)
  331. if self.device == torch.device("cpu"):
  332. counter_node = 0
  333. progress_bar = tqdm.tqdm(total = len(self.nodes), leave = False)
  334. progress_bar.set_description_str("[INFO]: Running PC backward pass")
  335. for level in self.traversal_order_backward[1:]:
  336. for node in level:
  337. progress_bar.n = counter_node + 1
  338. progress_bar.refresh()
  339. counter_node += 1
  340. node.backward()
  341. progress_bar.close()
  342. else:
  343. for level in self.traversal_order_backward[1:]:
  344. for node in level:
  345. node.backward()
  346. return
  347. def forward(self):
  348. if not self.reuse_forward:
  349. if len(self.traversal_order_forward) == 0:
  350. logger.log_fatal("Empty tree. Quit.")
  351. exit(-1)
  352. if self.root_node is None or self.traversal_order_forward[-1][0].id != self.root_node.id:
  353. logger.log_fatal("Missing root node. Quit.")
  354. exit(-1)
  355. self.reuse_backward = False
  356. self.reuse_forward = True
  357. if self.device == torch.device("cpu"):
  358. counter_node = 0
  359. progress_bar = tqdm.tqdm(total = len(self.nodes), leave = False)
  360. progress_bar.set_description_str("[INFO]: Running PC forward pass")
  361. for level in self.traversal_order_forward:
  362. for node in level:
  363. progress_bar.n = counter_node + 1
  364. progress_bar.refresh()
  365. counter_node += 1
  366. node.forward()
  367. progress_bar.close()
  368. else:
  369. for level in self.traversal_order_forward:
  370. for node in level:
  371. node.forward()
  372. return self.root_node.value_forward
  373. def gather_induced_trees(self):
  374. for induced_tree in self.recurse_induced_trees(self.root_node):
  375. induced_tree[1] = numpy.prod(induced_tree[1])
  376. self.induced_trees.append(induced_tree)
  377. return
  378. def get_weights(self):
  379. weights = []
  380. for sum_node in self.sum_nodes:
  381. weights.append(torch.clone(sum_node.weights))
  382. return weights
  383. def load(self, file_path_pc):
  384. if not os.path.exists(file_path_pc):
  385. logger.log_fatal("Invalid PC file path. Quit.")
  386. exit(-1)
  387. self.reuse_backward = False
  388. self.reuse_forward = False
  389. with open(file_path_pc, "r") as file_pc:
  390. categorical_leaf_node_id = -1
  391. counter_line = 0
  392. id_to_nodes = {}
  393. lines = file_pc.readlines()
  394. progress_bar = tqdm.tqdm(total = len(lines), leave = False)
  395. reading_nodes = True
  396. progress_bar.set_description_str("[INFO]: Loading PC")
  397. for line in lines:
  398. progress_bar.n = counter_line + 1
  399. progress_bar.refresh()
  400. counter_line += 1
  401. line = line.strip()
  402. if line[0] == "#":
  403. line = line.replace("#", "")
  404. if line == "NODES":
  405. reading_nodes = True
  406. elif line == "EDGES":
  407. reading_nodes = False
  408. continue
  409. line_list = line.split(",")
  410. if reading_nodes:
  411. node_id = int(line_list[0])
  412. node_type = line_list[1]
  413. if node_type == "SUM":
  414. sum_node = SumNode()
  415. sum_node.device = self.device
  416. sum_node.id = node_id
  417. self.nodes.append(sum_node)
  418. id_to_nodes[sum_node.id] = sum_node
  419. self.sum_nodes.append(sum_node)
  420. elif node_type == "PRD":
  421. product_node = ProductNode()
  422. product_node.device = self.device
  423. product_node.id = node_id
  424. self.nodes.append(product_node)
  425. id_to_nodes[product_node.id] = product_node
  426. self.product_nodes.append(product_node)
  427. elif node_type == "CatNode" or node_type == "CATNODE":
  428. categorical_leaf_node_list = []
  429. node_attribute_index = int(line_list[2])
  430. node_probabilities = line_list[3:]
  431. for i in range(0, len(node_probabilities)):
  432. node_probabilities[i] = float(node_probabilities[i])
  433. if node_attribute_index in self.leaf_nodes_dict.keys():
  434. categorical_leaf_node_list = self.leaf_nodes_dict[node_attribute_index]
  435. else:
  436. for category_index in range(0, len(node_probabilities)):
  437. categorical_leaf_node = CategoricalLeafNode()
  438. categorical_leaf_node.attribute_index = node_attribute_index
  439. categorical_leaf_node.category_index = category_index
  440. categorical_leaf_node.device = self.device
  441. categorical_leaf_node.id = categorical_leaf_node_id
  442. categorical_leaf_node_id -= 1
  443. categorical_leaf_node_list.append(categorical_leaf_node)
  444. self.leaf_nodes.append(categorical_leaf_node)
  445. self.nodes.append(categorical_leaf_node)
  446. id_to_nodes[categorical_leaf_node.id] = categorical_leaf_node
  447. self.leaf_nodes_dict[node_attribute_index] = categorical_leaf_node_list
  448. sum_node = SumNode()
  449. sum_node.children = categorical_leaf_node_list
  450. sum_node.device = self.device
  451. sum_node.id = node_id
  452. sum_node.leaf = True
  453. sum_node.weights = node_probabilities
  454. self.nodes.append(sum_node)
  455. id_to_nodes[sum_node.id] = sum_node
  456. self.sum_nodes.append(sum_node)
  457. for (i, categorical_leaf_node) in enumerate(categorical_leaf_node_list):
  458. sum_node.weights_index_by_child_id[categorical_leaf_node.id] = i
  459. categorical_leaf_node.parents.append(sum_node)
  460. elif node_type == "CATNODEPRD":
  461. node_attribute_index = int(line_list[2])
  462. node_category_index = int(line_list[3])
  463. categorical_leaf_node = CategoricalLeafNode()
  464. categorical_leaf_node.attribute_index = node_attribute_index
  465. categorical_leaf_node.category_index = node_category_index
  466. categorical_leaf_node.device = self.device
  467. categorical_leaf_node.id = node_id
  468. self.leaf_nodes.append(categorical_leaf_node)
  469. self.nodes.append(categorical_leaf_node)
  470. id_to_nodes[categorical_leaf_node.id] = categorical_leaf_node
  471. else:
  472. nodes = []
  473. node_id_first = int(line_list[0])
  474. node_id_second = int(line_list[1])
  475. nodes.append(id_to_nodes[node_id_first])
  476. nodes.append(id_to_nodes[node_id_second])
  477. if len(nodes) != 2:
  478. logger.log_fatal("Invalid edge. Quit.")
  479. exit(-1)
  480. if len(line_list) >= 3:
  481. node_weight = float(line_list[2])
  482. if isinstance(nodes[0], SumNode) and not nodes[0].leaf:
  483. nodes[0].children.append(nodes[1])
  484. nodes[0].weights.append(node_weight)
  485. nodes[0].weights_index_by_child_id[nodes[1].id] = len(nodes[0].weights) - 1
  486. nodes[1].parents.append(nodes[0])
  487. elif isinstance(nodes[1], SumNode) and not nodes[0].leaf:
  488. nodes[1].children.append(nodes[0])
  489. nodes[1].weights.append(node_weight)
  490. nodes[1].weights_index_by_child_id[nodes[0].id] = len(nodes[1].weights) - 1
  491. nodes[0].parents.append(nodes[1])
  492. else:
  493. if isinstance(nodes[0], ProductNode):
  494. nodes[0].children.append(nodes[1])
  495. nodes[1].parents.append(nodes[0])
  496. elif isinstance(nodes[1], ProductNode):
  497. nodes[1].children.append(nodes[0])
  498. nodes[0].parents.append(nodes[1])
  499. progress_bar.close()
  500. for sum_node in self.sum_nodes:
  501. sum_node.weights = torch.Tensor(sum_node.weights).reshape(-1, 1)
  502. sum_node.weights = sum_node.weights.to(self.device)
  503. root_nodes = []
  504. for node in self.nodes:
  505. if len(node.parents) == 0:
  506. root_nodes.append(node)
  507. if len(root_nodes) != 1:
  508. logger.log_fatal("Invalid PC. Quit.")
  509. exit(-1)
  510. self.depth = self.traverse({0: root_nodes}, 0)
  511. self.root_node = root_nodes[0]
  512. self.traversal_order_backward = self.topological_sort_backward()
  513. self.traversal_order_forward = self.traversal_order_backward.copy()
  514. self.traversal_order_forward.reverse()
  515. def getLengthNestedList(nested_list):
  516. length = 0
  517. for item in nested_list:
  518. if isinstance(item, list):
  519. length += getLengthNestedList(item)
  520. else:
  521. length += 1
  522. return length
  523. if getLengthNestedList(self.traversal_order_backward) != len(self.nodes):
  524. logger.log_fatal("Invalid PC backward traversal. Quit.")
  525. exit(-1)
  526. if getLengthNestedList(self.traversal_order_forward) != len(self.nodes):
  527. logger.log_fatal("Invalid PC forward traversal. Quit.")
  528. exit(-1)
  529. return
  530. def normalize_weights(self, smoothing_epsilon):
  531. for sum_node in self.sum_nodes:
  532. value_forward_children = []
  533. weight_normalization = 0
  534. for child in sum_node.children:
  535. value_forward_children.append(child.value_forward)
  536. value_forward_children = torch.stack(value_forward_children)
  537. value_forward_children_max = torch.max(value_forward_children, 0)[0]
  538. for (i, child) in enumerate(sum_node.children):
  539. weight_normalization += sum_node.weights[i] * torch.exp(child.value_forward - value_forward_children_max) + smoothing_epsilon
  540. # Local weight normalization with Laplace smoothing
  541. for (i, child) in enumerate(sum_node.children):
  542. weight = sum_node.weights[i] * torch.exp(child.value_forward - value_forward_children_max) + smoothing_epsilon
  543. sum_node.weights[i] = weight / weight_normalization
  544. return
  545. def randomize_weights(self):
  546. self.reuse_backward = False
  547. self.reuse_forward = False
  548. weights = self.get_weights()
  549. for i in range(len(weights)):
  550. weights[i] = weights[i].uniform_(0, 1)
  551. weights[i] /= torch.sum(weights[i])
  552. self.set_weights(weights)
  553. return
  554. def recurse_induced_trees(self, node):
  555. if isinstance(node, SumNode):
  556. for (i, child) in enumerate(node.children):
  557. for sub_induced_tree in self.recurse_induced_trees(child):
  558. induced_tree = [{node}, [node.weights[i].item()]]
  559. induced_tree[0].update(sub_induced_tree[0])
  560. induced_tree[1] += sub_induced_tree[1]
  561. yield induced_tree
  562. elif isinstance(node, ProductNode):
  563. sub_induced_trees_product = []
  564. for child in node.children:
  565. sub_induced_trees = []
  566. for sub_induced_tree in self.recurse_induced_trees(child):
  567. sub_induced_trees.append(sub_induced_tree)
  568. sub_induced_trees_product.append(sub_induced_trees)
  569. for sub_induced_trees in itertools.product(*sub_induced_trees_product):
  570. induced_tree = [{node}, []]
  571. for sub_induced_tree in sub_induced_trees:
  572. induced_tree[0].update(sub_induced_tree[0])
  573. induced_tree[1] += sub_induced_tree[1]
  574. yield induced_tree
  575. elif isinstance(node, CategoricalLeafNode):
  576. yield [{node}, []]
  577. return
  578. def set_leaf_nodes_binary(self, settings):
  579. self.reuse_backward = False
  580. self.reuse_forward = False
  581. self.batch_size = settings[0].shape[0]
  582. for leaf_node in self.leaf_nodes:
  583. leaf_node.set_binary(settings)
  584. return
  585. def set_leaf_nodes_categorical(self, settings):
  586. self.reuse_backward = False
  587. self.reuse_forward = False
  588. self.batch_size = settings.shape[0]
  589. for leaf_node in self.leaf_nodes:
  590. leaf_node.set_categorical(settings)
  591. return
  592. def set_weights(self, weights):
  593. if len(weights) != len(self.sum_nodes):
  594. logger.log_fatal("Invalid weights. Quit.")
  595. exit(-1)
  596. self.reuse_backward = False
  597. self.reuse_forward = False
  598. for (sum_node, sum_node_weights) in zip(self.sum_nodes, weights):
  599. sum_node.weights = torch.clone(sum_node_weights).to(self.device)
  600. return
  601. def topological_sort_backward(self):
  602. levels = []
  603. parent_count = {}
  604. queue = []
  605. for node in self.nodes:
  606. parent_count[node.id] = len(node.parents)
  607. queue.append(self.root_node)
  608. while len(queue) > 0:
  609. level = []
  610. queue_next_level = []
  611. for node in queue:
  612. level.append(node)
  613. for child in node.children:
  614. parent_count[child.id] -= 1
  615. if parent_count[child.id] <= 0:
  616. queue_next_level.append(child)
  617. levels.append(level)
  618. queue = queue_next_level.copy()
  619. return levels
  620. def traverse(self, layers, depth):
  621. if depth not in layers.keys():
  622. return depth - 1
  623. for node in layers[depth]:
  624. node.depth = depth
  625. for child in node.children:
  626. if depth + 1 not in layers.keys():
  627. layers[depth + 1] = []
  628. layers[depth + 1].append(child)
  629. depth_next = self.traverse(layers, depth + 1)
  630. if depth_next > depth:
  631. depth = depth_next
  632. return depth

pc.py at commit 4a46c28, under CC-BY-NC-SA-4.0 · at the source

Overview

Authors: Weixin Chen1, Simon Yu2, Huajie Shao3, Lui Sha1, Han Zhao1
  1. Department of Computer Science, University of Illinois Urbana-Champaign, Urbana, IL USA
  2. Department of Electrical and Computer Engineering, University of Illinois Urbana-Champaign, Urbana, IL USA
  3. Department of Computer Science, College of William and Mary, Williamsburg, VA USA
Institutions: University of Illinois Urbana-Champaign (United States); William & Mary (United States)
Journal: Machine learning, volume 115, issue 9, article 207
Dates: received 12 January 2025; accepted 2 July 2026; published online 29 August 2026; in print 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1007/s10994-026-07118-7 · PMID 42670317 · PMCID PMC13525958 · OpenAlex W4406692178
Open access: hybrid, a free copy (OpenAlex)
Status: code verified
Categories: systems (subfield)
Keywords: Interpretability, Compositional models, Probabilistic circuits, Logical reasoning
Topic: Neural Networks and Applications (Artificial Intelligence, Computer Science), according to OpenAlex
Citations: not cited yet (Europe PMC); 109 references in the paper

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

License: CC-BY-NC-SA-4.0
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 4a46c28f7cbfcc4fd89c5460de74c573fd14002c, 9 October 2025
Languages: Python (18), Shell (10)
Size: 36 files, 28 scripts
Software Heritage: not archived
Found in: “Code Availability”
Holds: README, license file, environment (requirements.txt), tests, documentation
Not found: CITATION.cff, continuous integration
Tools: PyTorch (14 files), NumPy (2 files), Pillow (1 file), scikit-learn (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
30 files

Code Availability

The code is available at https://github.com/uiuctml/npc-models.

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

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://cvml.ista.ac.at/AwA2/. CelebA is available at https://mmlab.ie.cuhk.edu.hk/projects/CelebA.html. GTSRB is availble at https://www.kaggle.com/datasets/meowmeowmeowmeowmeow/gtsrb-german-traffic-sign. MNIST is available at https://huggingface.co/datasets/ylecun/mnist. Three out of the four datasets, namely, AwA2, GTSRB, and MNIST, grants us rights to freely use, redistribute, and publish portions of the datasets for research purposes. CelebA, albeit granting limited access for non-commercial research purposes, explicitly forbids redistributions or publications of any portion of the dataset, as per its terms of use. Therefore, all instances of CelebA images are redacted from this work in compliance with such terms of use. During our experiments, we have generated additional data for all four datasets, such as annotations in regards to attribute recognition. Upon acceptance, all data generated during our experiments will be released and remain publicly available under appropriate licenses.

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://doi.org/10.1007/s10994-026-07118-7

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/s10994-026-07118-7},
url = {https://doi.org/10.1007/s10994-026-07118-7},
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/08/29
VL - 115
IS - 9
SP - 207
SN - 0885-6125
PB - Springer Science+Business Media
DO - 10.1007/s10994-026-07118-7
UR - https://doi.org/10.1007/s10994-026-07118-7
LA - en
ER -

CSL-JSON

{
"id": "10.1007/s10994-026-07118-7",
"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": "Mach Learn",
"volume": "115",
"issue": "9",
"page": "207",
"DOI": "10.1007/s10994-026-07118-7",
"PMID": "42670317",
"PMCID": "PMC13525958",
"ISSN": "0885-6125",
"publisher": "Springer Science+Business Media",
"URL": "https://doi.org/10.1007/s10994-026-07118-7",
"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 communications
In 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 insights
In 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: iScience
In 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 biology
In 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 neuroscience
In 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 Neuroscience
In 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: iScience
In 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 communications
In 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: iScience
In 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 biology
In 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.

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.