GitLab Repo

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