Skip to content

Commit edb1f84

Browse files
committed
More Podlang parser performance improvements
1 parent 7bea7db commit edb1f84

7 files changed

Lines changed: 170 additions & 67 deletions

File tree

Cargo.toml

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -37,6 +37,7 @@ hashbrown = { version = "0.14.3", default-features = false, features = [
3737
pest = "2.8.0"
3838
pest_derive = "2.8.0"
3939
petgraph = "0.6"
40+
rayon = "1.10"
4041
directories = { version = "6.0.0", optional = true }
4142
minicbor-serde = { version = "0.5.0", features = ["std"], optional = true }
4243
serde_bytes = "0.11"

src/frontend/custom.rs

Lines changed: 12 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,10 @@
11
#![allow(unused)]
2-
use std::{collections::HashMap, fmt, hash as h, iter, iter::zip, sync::Arc};
2+
use std::{
3+
collections::{HashMap, HashSet},
4+
fmt, hash as h, iter,
5+
iter::zip,
6+
sync::Arc,
7+
};
38

49
use schemars::JsonSchema;
510

@@ -132,6 +137,9 @@ pub struct CustomPredicateBatchBuilder {
132137
params: Params,
133138
pub name: String,
134139
pub predicates: Vec<CustomPredicate>,
140+
/// Names added via `predicate()`, kept alongside `predicates` so the
141+
/// duplicate-name check doesn't rescan the whole batch per insertion.
142+
predicate_names: HashSet<String>,
135143
/// Forward references to resolve in finish(): (predicate_idx, statement_idx, arg_idx, name)
136144
pending_self_pred_hashes: Vec<(usize, usize, usize, String)>,
137145
}
@@ -142,6 +150,7 @@ impl CustomPredicateBatchBuilder {
142150
params,
143151
name,
144152
predicates: Vec::new(),
153+
predicate_names: HashSet::new(),
145154
pending_self_pred_hashes: Vec::new(),
146155
}
147156
}
@@ -176,7 +185,7 @@ impl CustomPredicateBatchBuilder {
176185
priv_args: &[&str],
177186
sts: &[StatementTmplBuilder],
178187
) -> Result<Predicate> {
179-
if self.predicates.iter().any(|p| p.name == name) {
188+
if self.predicate_names.contains(name) {
180189
return Err(Error::custom(format!(
181190
"Duplicate predicate name '{}' in batch",
182191
name
@@ -271,6 +280,7 @@ impl CustomPredicateBatchBuilder {
271280
.collect(),
272281
)?;
273282
self.predicates.push(custom_predicate);
283+
self.predicate_names.insert(name.to_string());
274284
self.pending_self_pred_hashes.extend(pending);
275285
Ok(Predicate::BatchSelf(self.predicates.len() - 1))
276286
}

src/lang/frontend_ast_lower.rs

Lines changed: 9 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -539,7 +539,7 @@ impl<'a> Lowerer<'a> {
539539
&self,
540540
) -> Result<Vec<frontend_ast_split::SplitResult>, LoweringError> {
541541
let doc = self.validated.document();
542-
let predicates: Vec<CustomPredicateDef> = doc
542+
let mut predicates: Vec<CustomPredicateDef> = doc
543543
.items
544544
.iter()
545545
.filter_map(|item| match item {
@@ -548,23 +548,15 @@ impl<'a> Lowerer<'a> {
548548
})
549549
.collect();
550550

551-
// Apply splitting to each predicate as needed. The typed-key rewrite
552-
// happens before splitting so split chain pieces inherit `Index` keys
553-
// unchanged. The search cache is shared across the module: modules
554-
// routinely contain families of same-shape predicates, which then
555-
// pay for the ordering search once.
556-
let split_started = std::time::Instant::now();
557-
let mut search_cache = frontend_ast_split::SplitSearchCache::default();
558-
let mut split_results = Vec::new();
559-
for mut pred in predicates {
560-
self.rewrite_typed_dot_access(&mut pred);
561-
let result = frontend_ast_split::split_predicate_if_needed(
562-
pred,
563-
self.params,
564-
&mut search_cache,
565-
)?;
566-
split_results.push(result);
551+
// The typed-key rewrite happens before splitting so split chain
552+
// pieces inherit `Index` keys unchanged.
553+
for pred in &mut predicates {
554+
self.rewrite_typed_dot_access(pred);
567555
}
556+
557+
let split_started = std::time::Instant::now();
558+
let split_results =
559+
frontend_ast_split::split_predicates_if_needed(predicates, self.params)?;
568560
log::debug!(
569561
"predicate splitting: {:?} ({} predicates)",
570562
split_started.elapsed(),

src/lang/frontend_ast_split.rs

Lines changed: 111 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -97,8 +97,9 @@ pub(super) fn validate_predicate_is_splittable(
9797
}
9898

9999
/// Split a predicate into a chain if it exceeds statement limit. Callers
100-
/// splitting many predicates under one `Params` (module lowering) share one
101-
/// cache so same-shape predicates search only once.
100+
/// splitting many predicates under one `Params` share one cache so same-shape
101+
/// predicates search only once; whole-module callers should prefer
102+
/// [`split_predicates_if_needed`], which also parallelizes the searches.
102103
pub fn split_predicate_if_needed(
103104
pred: CustomPredicateDef,
104105
params: &Params,
@@ -124,6 +125,73 @@ pub fn split_predicate_if_needed(
124125
})
125126
}
126127

128+
/// Split a whole module's predicates. The link assignment is a pure function
129+
/// of a predicate's statement/wildcard shape, so each distinct shape is
130+
/// searched once and the searches run in parallel across shapes.
131+
pub fn split_predicates_if_needed(
132+
predicates: Vec<CustomPredicateDef>,
133+
params: &Params,
134+
) -> Result<Vec<SplitResult>, SplittingError> {
135+
use rayon::prelude::*;
136+
137+
for pred in &predicates {
138+
validate_predicate_is_splittable(pred)?;
139+
}
140+
141+
// Predicates over the statement cap get a prepared split input; the rest
142+
// pass through whole.
143+
let inputs: Vec<Option<SplitInput>> = predicates
144+
.iter()
145+
.map(|pred| {
146+
(pred.statements.len() > Params::max_custom_predicate_arity())
147+
.then(|| prepare_split_input(pred))
148+
})
149+
.collect();
150+
151+
// One representative per distinct shape, in first-encounter order so a
152+
// search failure surfaces the same predicate the sequential loop would
153+
// have blamed.
154+
let mut seen_shapes: HashSet<&SplitShape> = HashSet::new();
155+
let mut representatives: Vec<(&str, &SplitInput)> = Vec::new();
156+
for (pred, input) in predicates.iter().zip(&inputs) {
157+
if let Some(input) = input {
158+
if seen_shapes.insert(&input.shape) {
159+
representatives.push((pred.name.name.as_str(), input));
160+
}
161+
}
162+
}
163+
164+
let searched: Vec<Result<LinkAssignment, SplittingError>> = representatives
165+
.par_iter()
166+
.map(|(name, input)| search_link_assignment(name, input, params))
167+
.collect();
168+
let mut assignment_of_shape: HashMap<&SplitShape, LinkAssignment> = HashMap::new();
169+
for ((_, input), result) in representatives.iter().zip(searched) {
170+
assignment_of_shape.insert(&input.shape, result?);
171+
}
172+
173+
let results: Vec<Result<SplitResult, SplittingError>> = predicates
174+
.into_par_iter()
175+
.zip(&inputs)
176+
.map(|(pred, input)| match input {
177+
None => Ok(SplitResult {
178+
predicates: vec![pred],
179+
chain_info: None,
180+
}),
181+
Some(input) => {
182+
let assignment = assignment_of_shape[&input.shape].clone();
183+
let (chain_predicates, chain_info) =
184+
build_chain_from_assignment(&pred, input, assignment, params)?;
185+
Ok(SplitResult {
186+
predicates: chain_predicates,
187+
chain_info: Some(chain_info),
188+
})
189+
}
190+
})
191+
.collect();
192+
results.into_iter().collect()
193+
}
194+
127195
fn collect_wildcards_from_statement(stmt: &StatementTmpl) -> HashSet<String> {
128196
stmt.wildcard_names().map(str::to_string).collect()
129197
}
@@ -852,6 +920,7 @@ enum CandidateSearch {
852920
fn candidate_orderings(input: &SplitInput, params: &Params) -> CandidateSearch {
853921
use rand::{seq::SliceRandom, SeedableRng};
854922
use rand_chacha::ChaCha20Rng;
923+
use rayon::prelude::*;
855924

856925
let num_statements = input.shape.num_statements;
857926
let max_args = Params::max_statement_args();
@@ -882,26 +951,29 @@ fn candidate_orderings(input: &SplitInput, params: &Params) -> CandidateSearch {
882951
refinement_seeds.push(shuffled);
883952
}
884953

885-
let mut refined: Vec<(Vec<usize>, usize)> = Vec::new();
886-
let mut shortcut_attempted = false;
887-
for seed in refinement_seeds {
888-
let result = refine_ordering(
889-
seed,
890-
&input.shape.statements_using,
891-
&input.shape.is_original_public,
892-
max_args,
893-
REFINE_ITERATIONS,
894-
);
895-
let result_cost = input.excess_cost(&result);
896-
if result_cost == 0 && !shortcut_attempted {
897-
shortcut_attempted = true;
898-
if let Some(assignment) = input.partition_at_lower_bound(&result, params) {
899-
return CandidateSearch::Found(assignment);
900-
}
901-
// A cost-0 ordering that still needs extra links: refine the
902-
// remaining seeds and let the full partition loop decide.
954+
// Each refinement depends only on its seed, so the seeds run in parallel.
955+
let mut refined: Vec<(Vec<usize>, usize)> = refinement_seeds
956+
.into_par_iter()
957+
.map(|seed| {
958+
let result = refine_ordering(
959+
seed,
960+
&input.shape.statements_using,
961+
&input.shape.is_original_public,
962+
max_args,
963+
REFINE_ITERATIONS,
964+
);
965+
let result_cost = input.excess_cost(&result);
966+
(result, result_cost)
967+
})
968+
.collect();
969+
// The first cost-0 refinement in seed order gets one shortcut attempt: the
970+
// stable sort below keeps seed order among ties, so it is the ordering the
971+
// partition loop would try first anyway. A cost-0 ordering that still
972+
// needs extra links falls through to the full partition loop.
973+
if let Some((zero_cost_ordering, _)) = refined.iter().find(|(_, cost)| *cost == 0) {
974+
if let Some(assignment) = input.partition_at_lower_bound(zero_cost_ordering, params) {
975+
return CandidateSearch::Found(assignment);
903976
}
904-
refined.push((result, result_cost));
905977
}
906978
refined.sort_by_key(|(_, refined_cost)| *refined_cost);
907979

@@ -931,12 +1003,7 @@ fn split_into_chain(
9311003
params: &Params,
9321004
search_cache: &mut SplitSearchCache,
9331005
) -> Result<(Vec<CustomPredicateDef>, SplitChainInfo), SplittingError> {
934-
let original_name = pred.name.name.clone();
935-
let conjunction = pred.conjunction_type;
936-
let real_statement_count = pred.statements.len();
937-
9381006
let input = prepare_split_input(pred);
939-
let num_statements = input.shape.num_statements;
9401007

9411008
// The link assignment depends only on the statement/wildcard usage shape
9421009
// (plus `params`, fixed for the cache's lifetime), and modules routinely
@@ -945,13 +1012,30 @@ fn split_into_chain(
9451012
let assignment = if let Some(assignment) = search_cache.assignments.get(&input.shape) {
9461013
assignment.clone()
9471014
} else {
948-
let assignment = search_link_assignment(&original_name, &input, params)?;
1015+
let assignment = search_link_assignment(&pred.name.name, &input, params)?;
9491016
search_cache
9501017
.assignments
9511018
.insert(input.shape.clone(), assignment.clone());
9521019
assignment
9531020
};
9541021

1022+
build_chain_from_assignment(pred, &input, assignment, params)
1023+
}
1024+
1025+
/// Everything downstream of the ordering search: reorder the statements per
1026+
/// the assignment, derive each link's argument lists, and emit the chain's
1027+
/// predicate definitions plus [`SplitChainInfo`].
1028+
fn build_chain_from_assignment(
1029+
pred: &CustomPredicateDef,
1030+
input: &SplitInput,
1031+
assignment: LinkAssignment,
1032+
params: &Params,
1033+
) -> Result<(Vec<CustomPredicateDef>, SplitChainInfo), SplittingError> {
1034+
let original_name = pred.name.name.clone();
1035+
let conjunction = pred.conjunction_type;
1036+
let real_statement_count = pred.statements.len();
1037+
let num_statements = input.shape.num_statements;
1038+
9551039
// Reorder map: original index -> position in flattened chain.
9561040
let mut reorder_map = vec![0usize; num_statements];
9571041
{

src/lang/frontend_ast_validate.rs

Lines changed: 12 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -540,13 +540,17 @@ impl Validator {
540540
pred_def: &CustomPredicateDef,
541541
) -> Result<(), ValidationError> {
542542
let pred_name = pred_def.name.name.clone();
543+
let wildcard_scope = self
544+
.symbols
545+
.wildcard_scopes
546+
.get(&pred_name)
547+
.expect("Wildcard scope should exist after pass 1");
548+
549+
// The collision check depends only on the scope, so it runs once per
550+
// predicate rather than once per statement.
551+
self.validate_wildcard_names(&wildcard_scope.wildcards.keys().collect())?;
543552

544553
for stmt in &pred_def.statements {
545-
let wildcard_scope = self
546-
.symbols
547-
.wildcard_scopes
548-
.get(&pred_name)
549-
.expect("Wildcard scope should exist after pass 1");
550554
self.validate_statement(stmt, Some((&pred_name, wildcard_scope)))?;
551555
}
552556

@@ -592,12 +596,6 @@ impl Validator {
592596
let pred_name = stmt.predicate.predicate_name();
593597
let pred_span = stmt.predicate.span();
594598

595-
let wc_names = match wildcard_context {
596-
Some((_, wc_scope)) => wc_scope.wildcards.keys().collect(),
597-
None => HashSet::new(),
598-
};
599-
self.validate_wildcard_names(&wc_names)?;
600-
601599
// Check if predicate exists
602600
let pred_info = match &stmt.predicate {
603601
PredicateRef::Qualified { module, predicate } => {
@@ -642,7 +640,9 @@ impl Validator {
642640
} else if let Some(info) = self.symbols.predicates.get(pred_name) {
643641
// Custom or imported predicate
644642
Some(info.clone())
645-
} else if wc_names.contains(&pred_name.to_string()) {
643+
} else if wildcard_context
644+
.is_some_and(|(_, scope)| scope.wildcards.contains_key(pred_name))
645+
{
646646
None
647647
} else {
648648
return Err(ValidationError::UndefinedPredicate {

src/lang/module.rs

Lines changed: 23 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -333,9 +333,31 @@ fn build_single_batch(
333333
params: &Params,
334334
batch_name: &str,
335335
) -> Result<Arc<CustomPredicateBatch>, BatchingError> {
336+
use rayon::prelude::*;
337+
336338
let mut builder = CustomPredicateBatchBuilder::new(params.clone(), batch_name.to_string());
337339

338-
for pred in predicates {
340+
// Statement-template construction is a pure function of the symbol table
341+
// and reference map, so it runs in parallel; only the batch insertion
342+
// below stays sequential.
343+
let prepared: Vec<Vec<StatementTmplBuilder>> = predicates
344+
.par_iter()
345+
.map(|pred| {
346+
pred.statements
347+
.iter()
348+
.map(|stmt| {
349+
build_statement_with_resolved_refs(
350+
stmt,
351+
reference_map,
352+
&pred.name.name,
353+
symbols,
354+
)
355+
})
356+
.collect::<Result<_, _>>()
357+
})
358+
.collect::<Result<_, _>>()?;
359+
360+
for (pred, statement_builders) in predicates.iter().zip(prepared) {
339361
let name = &pred.name.name;
340362

341363
// Collect argument names
@@ -353,13 +375,6 @@ fn build_single_batch(
353375
.map(|args| args.iter().map(|a| a.name.as_str()).collect())
354376
.unwrap_or_default();
355377

356-
// Build statement templates with resolved predicates
357-
let statement_builders: Vec<StatementTmplBuilder> = pred
358-
.statements
359-
.iter()
360-
.map(|stmt| build_statement_with_resolved_refs(stmt, reference_map, name, symbols))
361-
.collect::<Result<_, _>>()?;
362-
363378
let conjunction = pred.conjunction_type == ConjunctionType::And;
364379

365380
builder

src/middleware/custom.rs

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -466,8 +466,9 @@ impl fmt::Debug for CustomPredicateBatchData {
466466
// TODO: Rename Batch for Module everywhere in the code base
467467
impl CustomPredicateBatchData {
468468
fn new_full(predicates: Vec<CustomPredicate>) -> Self {
469+
use rayon::prelude::*;
469470
let kvs: HashMap<RawValue, RawValue> = predicates
470-
.iter()
471+
.par_iter()
471472
.enumerate()
472473
.map(|(index, pred)| {
473474
let cp_hash = hash_fields(&pred.to_fields());

0 commit comments

Comments
 (0)