amachine.am_optimize.am_optimize_symbol_permutation_cpsat
1import numpy as np 2import warnings 3from collections import defaultdict, Counter 4 5from amachine.am_random import resolve_rng 6 7def optimize_symbol_permutation_cpsat( 8 m, 9 symbol_set, 10 frequency_bias, 11 n_search_workers, 12 max_search_time, 13 np_rng : np.random.Generator | None = None ): 14 15 from ortools.sat.python import cp_model 16 17 np_rng = resolve_rng( np_rng ) 18 random_seed = np_rng.integers(0, 2**31 - 1) 19 20 transitions = m.transitions 21 n_edges = len(transitions) 22 23 print(f"Random machine has {n_edges} transitions.") 24 25 state_trs_in = defaultdict(list) 26 state_trs_out = defaultdict(list) 27 for i, tr in enumerate(transitions): 28 state_trs_in[ tr.target_state_idx ].append((i, tr)) 29 state_trs_out[ tr.origin_state_idx ].append((i, tr)) 30 31 pi = m.get_stationary_distribution() 32 tr_pair_probs = {} 33 for state_idx in range(len(m.states)): 34 for i_in, t_in in state_trs_in[state_idx]: 35 in_prob = pi[t_in.origin_state_idx] * t_in.prob 36 for i_out, t_out in state_trs_out[state_idx]: 37 tr_pair_probs[(i_in, i_out)] = in_prob * t_out.prob 38 39 n_pairs = len(tr_pair_probs) 40 uniform_w = 1.0 / n_pairs if n_pairs > 0 else 0.0 41 PRECISION_SCALE = 10_000_000 42 43 # Pre-filter zero-weight pairs before allocating any solver variables. 44 weighted_pairs = {} 45 for pair, p in tr_pair_probs.items(): 46 coeff = int(((1.0 - frequency_bias) * uniform_w + frequency_bias * p) * PRECISION_SCALE) 47 if coeff > 0: 48 weighted_pairs[pair] = coeff 49 50 model = cp_model.CpModel() 51 52 original_symbols = [t.symbol_idx for t in transitions] 53 symbol_counts = Counter(original_symbols) 54 present_symbols = sorted(symbol_counts.keys()) 55 n_present = len(present_symbols) 56 57 # remap the n_present used symbols to consecutive indices 0..n_present-1. 58 # This makes sym[i]'s domain contiguous and shrinks the idx variable's 59 # range from [0, n_symbols^2-1] to [0, n_present^2-1], giving the LP relaxation 60 # much tighter bounds on every element-constraint pair. 61 k_to_sym = present_symbols # k -> original symbol index 62 sym_to_k = {s: k for k, s in enumerate(present_symbols)} 63 original_k = [sym_to_k[s] for s in original_symbols] # remapped original assignment 64 65 sym = [model.NewIntVar(0, n_present - 1, f"sym_{i}") for i in range(n_edges)] 66 for i in range(n_edges): 67 model.AddHint(sym[i], original_k[i]) 68 69 # Indicator booleans — keep the original OnlyEnforceIf formulation. 70 # These create native Boolean implication clauses inside CP-SAT's SAT engine, 71 # which the WeightedSum-equality replacement from the previous attempt destroyed. 72 b = [[model.NewBoolVar(f"b_{i}_{k}") for k in range(n_present)] 73 for i in range(n_edges)] 74 75 for i in range(n_edges): 76 # AddExactlyOne gives the SAT engine an explicit partition clause, 77 # complementing the OnlyEnforceIf implications in both directions. 78 model.AddExactlyOne(b[i]) 79 for k in range(n_present): 80 model.Add(sym[i] == k).OnlyEnforceIf(b[i][k]) 81 model.Add(sym[i] != k).OnlyEnforceIf(b[i][k].Not()) 82 model.AddHint(b[i][k], int(k == original_k[i])) 83 84 # Count constraints: column sums of the b matrix. 85 for k in range(n_present): 86 model.Add( 87 cp_model.LinearExpr.Sum([b[i][k] for i in range(n_edges)]) 88 == symbol_counts[k_to_sym[k]] 89 ) 90 91 # Unifilarity: skip the trivial single-edge-state case. 92 for out_edges in state_trs_out.values(): 93 if len(out_edges) > 1: 94 model.AddAllDifferent([sym[i] for i, _ in out_edges]) 95 96 # Remapped cohesion table: flat_coh[k_in * n_present + k_out]. 97 # n_present^2 entries instead of n_symbols^2 — smaller table, tighter idx domain. 98 coh_int = np.round(symbol_set.symbol_cohesion * PRECISION_SCALE).astype(int) 99 flat_coh = [int(coh_int[k_to_sym[k_in], k_to_sym[k_out]]) 100 for k_in in range(n_present) 101 for k_out in range(n_present)] 102 103 r_domain = cp_model.Domain.FromValues(sorted(set(flat_coh))) 104 idx_max = n_present * n_present - 1 105 106 reward_vars: list = [] 107 coeffs: list = [] 108 109 for (i_in, i_out), coeff in weighted_pairs.items(): 110 # idx now lives in [0, n_present^2-1] instead of [0, n_symbols^2-1]. 111 idx = model.NewIntVar(0, idx_max, f"idx_{i_in}_{i_out}") 112 model.Add(idx == sym[i_in] * n_present + sym[i_out]) 113 r = model.NewIntVarFromDomain(r_domain, f"r_{i_in}_{i_out}") 114 model.AddElement(idx, flat_coh, r) 115 orig_idx = original_k[i_in] * n_present + original_k[i_out] 116 model.AddHint(idx, orig_idx) 117 model.AddHint(r, flat_coh[orig_idx]) 118 reward_vars.append(r) 119 coeffs.append(coeff) 120 121 if reward_vars: 122 model.Maximize(cp_model.LinearExpr.WeightedSum(reward_vars, coeffs)) 123 124 solver = cp_model.CpSolver() 125 # solver.parameters.log_search_progress = True 126 solver.parameters.num_search_workers = n_search_workers 127 solver.parameters.max_time_in_seconds = max_search_time 128 solver.parameters.linearization_level = 0 129 solver.parameters.random_seed = int(random_seed) 130 131 res = [tr.symbol_idx for tr in m.transitions] 132 133 status = solver.Solve(model) 134 if status in (cp_model.OPTIMAL, cp_model.FEASIBLE): 135 msg = "OPTIMAL" if status == cp_model.OPTIMAL else "FEASIBLE" 136 print(f"{msg} solution found. Objective = {solver.ObjectiveValue():.0f}") 137 for i in range(n_edges): 138 res[i] = k_to_sym[solver.Value(sym[i])] # remap k back to original symbol index 139 else: 140 warnings.warn("optimization found no solution within time limit.") 141 142 return res
def
optimize_symbol_permutation_cpsat( m, symbol_set, frequency_bias, n_search_workers, max_search_time, np_rng: numpy.random._generator.Generator | None = None):
8def optimize_symbol_permutation_cpsat( 9 m, 10 symbol_set, 11 frequency_bias, 12 n_search_workers, 13 max_search_time, 14 np_rng : np.random.Generator | None = None ): 15 16 from ortools.sat.python import cp_model 17 18 np_rng = resolve_rng( np_rng ) 19 random_seed = np_rng.integers(0, 2**31 - 1) 20 21 transitions = m.transitions 22 n_edges = len(transitions) 23 24 print(f"Random machine has {n_edges} transitions.") 25 26 state_trs_in = defaultdict(list) 27 state_trs_out = defaultdict(list) 28 for i, tr in enumerate(transitions): 29 state_trs_in[ tr.target_state_idx ].append((i, tr)) 30 state_trs_out[ tr.origin_state_idx ].append((i, tr)) 31 32 pi = m.get_stationary_distribution() 33 tr_pair_probs = {} 34 for state_idx in range(len(m.states)): 35 for i_in, t_in in state_trs_in[state_idx]: 36 in_prob = pi[t_in.origin_state_idx] * t_in.prob 37 for i_out, t_out in state_trs_out[state_idx]: 38 tr_pair_probs[(i_in, i_out)] = in_prob * t_out.prob 39 40 n_pairs = len(tr_pair_probs) 41 uniform_w = 1.0 / n_pairs if n_pairs > 0 else 0.0 42 PRECISION_SCALE = 10_000_000 43 44 # Pre-filter zero-weight pairs before allocating any solver variables. 45 weighted_pairs = {} 46 for pair, p in tr_pair_probs.items(): 47 coeff = int(((1.0 - frequency_bias) * uniform_w + frequency_bias * p) * PRECISION_SCALE) 48 if coeff > 0: 49 weighted_pairs[pair] = coeff 50 51 model = cp_model.CpModel() 52 53 original_symbols = [t.symbol_idx for t in transitions] 54 symbol_counts = Counter(original_symbols) 55 present_symbols = sorted(symbol_counts.keys()) 56 n_present = len(present_symbols) 57 58 # remap the n_present used symbols to consecutive indices 0..n_present-1. 59 # This makes sym[i]'s domain contiguous and shrinks the idx variable's 60 # range from [0, n_symbols^2-1] to [0, n_present^2-1], giving the LP relaxation 61 # much tighter bounds on every element-constraint pair. 62 k_to_sym = present_symbols # k -> original symbol index 63 sym_to_k = {s: k for k, s in enumerate(present_symbols)} 64 original_k = [sym_to_k[s] for s in original_symbols] # remapped original assignment 65 66 sym = [model.NewIntVar(0, n_present - 1, f"sym_{i}") for i in range(n_edges)] 67 for i in range(n_edges): 68 model.AddHint(sym[i], original_k[i]) 69 70 # Indicator booleans — keep the original OnlyEnforceIf formulation. 71 # These create native Boolean implication clauses inside CP-SAT's SAT engine, 72 # which the WeightedSum-equality replacement from the previous attempt destroyed. 73 b = [[model.NewBoolVar(f"b_{i}_{k}") for k in range(n_present)] 74 for i in range(n_edges)] 75 76 for i in range(n_edges): 77 # AddExactlyOne gives the SAT engine an explicit partition clause, 78 # complementing the OnlyEnforceIf implications in both directions. 79 model.AddExactlyOne(b[i]) 80 for k in range(n_present): 81 model.Add(sym[i] == k).OnlyEnforceIf(b[i][k]) 82 model.Add(sym[i] != k).OnlyEnforceIf(b[i][k].Not()) 83 model.AddHint(b[i][k], int(k == original_k[i])) 84 85 # Count constraints: column sums of the b matrix. 86 for k in range(n_present): 87 model.Add( 88 cp_model.LinearExpr.Sum([b[i][k] for i in range(n_edges)]) 89 == symbol_counts[k_to_sym[k]] 90 ) 91 92 # Unifilarity: skip the trivial single-edge-state case. 93 for out_edges in state_trs_out.values(): 94 if len(out_edges) > 1: 95 model.AddAllDifferent([sym[i] for i, _ in out_edges]) 96 97 # Remapped cohesion table: flat_coh[k_in * n_present + k_out]. 98 # n_present^2 entries instead of n_symbols^2 — smaller table, tighter idx domain. 99 coh_int = np.round(symbol_set.symbol_cohesion * PRECISION_SCALE).astype(int) 100 flat_coh = [int(coh_int[k_to_sym[k_in], k_to_sym[k_out]]) 101 for k_in in range(n_present) 102 for k_out in range(n_present)] 103 104 r_domain = cp_model.Domain.FromValues(sorted(set(flat_coh))) 105 idx_max = n_present * n_present - 1 106 107 reward_vars: list = [] 108 coeffs: list = [] 109 110 for (i_in, i_out), coeff in weighted_pairs.items(): 111 # idx now lives in [0, n_present^2-1] instead of [0, n_symbols^2-1]. 112 idx = model.NewIntVar(0, idx_max, f"idx_{i_in}_{i_out}") 113 model.Add(idx == sym[i_in] * n_present + sym[i_out]) 114 r = model.NewIntVarFromDomain(r_domain, f"r_{i_in}_{i_out}") 115 model.AddElement(idx, flat_coh, r) 116 orig_idx = original_k[i_in] * n_present + original_k[i_out] 117 model.AddHint(idx, orig_idx) 118 model.AddHint(r, flat_coh[orig_idx]) 119 reward_vars.append(r) 120 coeffs.append(coeff) 121 122 if reward_vars: 123 model.Maximize(cp_model.LinearExpr.WeightedSum(reward_vars, coeffs)) 124 125 solver = cp_model.CpSolver() 126 # solver.parameters.log_search_progress = True 127 solver.parameters.num_search_workers = n_search_workers 128 solver.parameters.max_time_in_seconds = max_search_time 129 solver.parameters.linearization_level = 0 130 solver.parameters.random_seed = int(random_seed) 131 132 res = [tr.symbol_idx for tr in m.transitions] 133 134 status = solver.Solve(model) 135 if status in (cp_model.OPTIMAL, cp_model.FEASIBLE): 136 msg = "OPTIMAL" if status == cp_model.OPTIMAL else "FEASIBLE" 137 print(f"{msg} solution found. Objective = {solver.ObjectiveValue():.0f}") 138 for i in range(n_edges): 139 res[i] = k_to_sym[solver.Value(sym[i])] # remap k back to original symbol index 140 else: 141 warnings.warn("optimization found no solution within time limit.") 142 143 return res