SearchLibrium 0.0.152__tar.gz → 0.0.153__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (64) hide show
  1. {searchlibrium-0.0.152/src/SearchLibrium.egg-info → searchlibrium-0.0.153}/PKG-INFO +1 -1
  2. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/pyproject.toml +1 -1
  3. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/MixedLogit.py +84 -94
  4. searchlibrium-0.0.153/src/SearchLibrium/MixedLogitGSE.py +221 -0
  5. searchlibrium-0.0.153/src/SearchLibrium/version.txt +1 -0
  6. {searchlibrium-0.0.152 → searchlibrium-0.0.153/src/SearchLibrium.egg-info}/PKG-INFO +1 -1
  7. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium.egg-info/SOURCES.txt +5 -1
  8. searchlibrium-0.0.153/src/data_generator/demonstrate_generation.py +85 -0
  9. searchlibrium-0.0.153/src/data_generator/fit_recovery.py +465 -0
  10. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/data_generator/latent_class_gen.py +76 -13
  11. searchlibrium-0.0.153/src/data_generator/print_synthetic_data.py +122 -0
  12. searchlibrium-0.0.152/src/SearchLibrium/version.txt +0 -1
  13. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/MANIFEST.in +0 -0
  14. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/README.md +0 -0
  15. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/setup.cfg +0 -0
  16. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/BEHier.py +0 -0
  17. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/Halton.py +0 -0
  18. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/Mode_Activity_Nested.py +0 -0
  19. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/RandomP.py +0 -0
  20. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/SEARCH_SM_MARIO.py +0 -0
  21. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/Two_Level_Nest.py +0 -0
  22. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/__init__.py +0 -0
  23. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/__main__.py +0 -0
  24. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/_choice_model.py +0 -0
  25. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/_device.py +0 -0
  26. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/banditsa.py +0 -0
  27. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/bhhh/minimize.py +0 -0
  28. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/boxcox_functions.py +0 -0
  29. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/call_meta.py +0 -0
  30. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/constraints_builder.py +0 -0
  31. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/data/Swissmetro_final.csv +0 -0
  32. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/data/TravelMode.csv +0 -0
  33. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/data/electricity.csv +0 -0
  34. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/harmony.py +0 -0
  35. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/hspbil.py +0 -0
  36. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/latent_class.py +0 -0
  37. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/main.py +0 -0
  38. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/main_debug.py +0 -0
  39. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/mdcev.py +0 -0
  40. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/misc.py +0 -0
  41. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/mixed_logit.py +0 -0
  42. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/mixed_nested.py +0 -0
  43. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/mixedrrm.py +0 -0
  44. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/multinomial_logit.py +0 -0
  45. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/multinomial_nested.py +0 -0
  46. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/multinomial_probit.py +0 -0
  47. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/ordered_logit.py +0 -0
  48. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/ordered_logit_mixed.py +0 -0
  49. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/predict.py +0 -0
  50. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/rrm.py +0 -0
  51. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/sample_data.py +0 -0
  52. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/sapbil.py +0 -0
  53. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/search.py +0 -0
  54. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/selection_models.py +0 -0
  55. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/setup.py +0 -0
  56. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/siman.py +0 -0
  57. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/test_lc_de.py +0 -0
  58. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/test_mario_searches.py +0 -0
  59. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/test_sapbil_vs_banditsa.py +0 -0
  60. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium/threshold.py +0 -0
  61. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium.egg-info/dependency_links.txt +0 -0
  62. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium.egg-info/entry_points.txt +0 -0
  63. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium.egg-info/requires.txt +0 -0
  64. {searchlibrium-0.0.152 → searchlibrium-0.0.153}/src/SearchLibrium.egg-info/top_level.txt +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: SearchLibrium
3
- Version: 0.0.152
3
+ Version: 0.0.153
4
4
  Summary: A Python package for econometric models driven by search
5
5
  Author: Mario Fernando Taco Morales, Alexander Paz, Prithvi Beeramole, Robert Burdett
6
6
  Author-email: Zeke Ahern <z.ahern@qut.edu.au>
@@ -82,7 +82,7 @@ Documentation = "https://pypi.org/project/SearchLibrium/"
82
82
  realpython = "SearchLibrium.__main__:main"
83
83
 
84
84
  [tool.bumpver]
85
- current_version = "0.0.152"
85
+ current_version = "0.0.153"
86
86
  version_pattern = "MAJOR.MINOR.PATCH"
87
87
  commit_message = "[skip ci] Bump version {old_version} -> {new_version}"
88
88
  commit = true
@@ -134,14 +134,13 @@ class MixedLogit(DiscreteChoiceModel):
134
134
  def setup(self, X, y, varnames=None, alts=None, isvars=None, transvars=None,
135
135
  transformation="boxcox", ids=None, weights=None, avail=None,
136
136
  randvars=None, panels=None, base_alt=None, fit_intercept=False,
137
- init_coeff=None, maxiter=2000, correlated_vars=None,
137
+ init_coeff=None, maxiter=1000, correlated_vars=None,
138
138
  n_draws=1000, halton=True, minimise_func=None,
139
- batch_size=None, halton_opts=None, ftol=1e-6,
139
+ batch_size=None, halton_opts=None, ftol=1e-12,
140
140
  gtol=1e-6, return_hess=True, return_grad=True, method="slsqp",
141
141
  save_fitted_params=True, mnl_init=True,
142
142
  de_init=False, de_popsize=4, de_maxiter=3, de_tol=0.5,
143
- de_polish=False, l1_penalty=0.1, l2_penalty=0.5,
144
- sd_penalty=0.001):
143
+ de_polish=False, l1_penalty=0.0):
145
144
  # {
146
145
  self.fit_intercept = fit_intercept
147
146
  # ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
@@ -187,8 +186,6 @@ class MixedLogit(DiscreteChoiceModel):
187
186
  self.total_fun_eval = 0
188
187
  self.method = method.lower() if hasattr(method, 'lower') else method
189
188
  self.l1_penalty = float(l1_penalty)
190
- self.l2_penalty = float(l2_penalty)
191
- self.sd_penalty = float(sd_penalty)
192
189
  self.jac = self.return_grad # scipy optimize parameter
193
190
  self.n_draws = n_draws
194
191
  self.batch_size = min(n_draws, batch_size) if batch_size is not None else n_draws
@@ -387,7 +384,9 @@ class MixedLogit(DiscreteChoiceModel):
387
384
  # 2x Kftrans - mean and lambda, 3x Krtrans - mean, s.d., lambda
388
385
  # Kchol, Kbw - relate to random variables, non-transformed
389
386
  # Kchol - cholesky matrix, Kbw the s.d. for random vars
390
- n_coeff = self.Kf + self.Kr + self.Kchol + self.Kbw + 2 * self.Kftrans + 3 * self.Krtrans
387
+ n_coeff = (self.Kf + self.Kr + self.Kchol + self.Kbw
388
+ + 2 * self.Kftrans + 3 * self.Krtrans
389
+ + self._n_coeff_extra())
391
390
 
392
391
  # Initalise coefficients using a multinomial logit model
393
392
  if self.mnl_init and self.init_coeff is None:
@@ -411,30 +410,38 @@ class MixedLogit(DiscreteChoiceModel):
411
410
  fit_intercept=False)
412
411
  mnl.fit()
413
412
 
414
- # mnl estimates -> mxl needs to add stdev to random variables
415
- self.init_coeff = mnl.coeff_est
416
-
417
- lower = self.Kf + 2 * self.Kftrans + self.Kr
418
- upper = lower + self.Krtrans
419
-
420
- # ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
421
- # ERROR HANDLING (FIX) - WHEN LESS COEFFICIENTS THAN "kf + kr"
422
- if lower > len(self.init_coeff):
423
- # {
424
- additional_elements_needed = lower - len(self.init_coeff)
425
- extra = np.full(additional_elements_needed, 0.1)
426
- self.init_coeff = np.concatenate((self.init_coeff, extra))
427
- # }
428
- # ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
429
-
430
- arr = self.init_coeff[:lower]
431
- # Use simple initialization matching searchlogit for better convergence
413
+ # Build init_coeff from MNL estimates, mapping each variable
414
+ # to its correct MXL slot (Bf for fixed, Br_b for random).
415
+ bf_init = np.zeros(self.Kf)
416
+ br_b_init = np.zeros(self.Kr)
417
+ bf_pos = br_pos = 0
418
+ for ii in range(len(self.varnames)):
419
+ if not mnl.fxidx[ii]:
420
+ continue
421
+ mnl_idx = int(np.sum(mnl.fxidx[:ii]))
422
+ mnl_val = mnl.coeff_est[mnl_idx] if mnl_idx < len(mnl.coeff_est) else 0.1
423
+ if self.fxidx[ii]:
424
+ bf_init[bf_pos] = mnl_val
425
+ bf_pos += 1
426
+ elif self.rvidx[ii]:
427
+ br_b_init[br_pos] = mnl_val
428
+ br_pos += 1
429
+
430
+ mnl_Kf = int(np.sum(mnl.fxidx))
431
+ mnl_Kftrans = int(np.sum(mnl.fxtransidx))
432
+ bftrans_b = mnl.coeff_est[mnl_Kf:mnl_Kf + mnl_Kftrans] if mnl_Kftrans > 0 else np.array([])
433
+ bftrans_l = mnl.coeff_est[mnl_Kf + mnl_Kftrans:mnl_Kf + 2 * mnl_Kftrans] if mnl_Kftrans > 0 else np.array([])
434
+
435
+ arr = np.concatenate([bf_init, br_b_init] + list(self._init_pad_arrays()) + [bftrans_b, bftrans_l])
432
436
  rep = np.repeat(0.1, self.Kchol + self.Kbw)
433
- self.init_coeff = np.concatenate((arr, rep, self.init_coeff[lower:upper],))
437
+ self.init_coeff = np.concatenate([arr, rep])
434
438
 
435
439
  if self.Krtrans:
436
- rep = np.repeat(0.1, self.Krtrans)
437
- self.init_coeff = np.concatenate((self.init_coeff, rep, self.init_coeff[-self.Krtrans:]))
440
+ self.init_coeff = np.concatenate([
441
+ self.init_coeff,
442
+ np.repeat(0.1, self.Krtrans),
443
+ np.repeat(0.1, self.Krtrans),
444
+ ])
438
445
  # }
439
446
 
440
447
  betas = np.repeat(0.1, n_coeff) if self.init_coeff is None else self.init_coeff
@@ -454,7 +461,8 @@ class MixedLogit(DiscreteChoiceModel):
454
461
  "flmbda": (lmda_bound, self.Kftrans),
455
462
  "br_trans_b": (any_bound, self.Krtrans),
456
463
  "br_trans_w": (any_bound, self.Krtrans),
457
- "rlmbda": (lmda_bound, self.Krtrans)
464
+ "rlmbda": (lmda_bound, self.Krtrans),
465
+ **self._bound_extra(),
458
466
  }
459
467
  bnds = [[bound[1][0]] * bound[1][1] for bound in bound_dict.items() if bound[1][1] > 0]
460
468
  bnds = list(itertools.chain.from_iterable(bnds))
@@ -553,7 +561,8 @@ class MixedLogit(DiscreteChoiceModel):
553
561
  "flmbda": (lmda_bound, self.Kftrans),
554
562
  "br_trans_b": (any_bound, self.Krtrans),
555
563
  "br_trans_w": (any_bound, self.Krtrans),
556
- "rlmbda": (lmda_bound, self.Krtrans)
564
+ "rlmbda": (lmda_bound, self.Krtrans),
565
+ **self._bound_extra(),
557
566
  }
558
567
 
559
568
  # This code makes a specific number of copies of each range
@@ -570,10 +579,13 @@ class MixedLogit(DiscreteChoiceModel):
570
579
  if getattr(self, '_jax', False):
571
580
  jax_result = self.optimize_jax(betas, draws, drawstrans)
572
581
  if jax_result is not None:
573
- beta_segment_names = ["Bf", "Br_b", "chol", "Br_w", "Bftrans",
574
- "flmbda", "Brtrans_b", "Brtrans_w", "rlmda"]
575
- iterations = [self.Kf, self.Kr, self.Kchol, self.Kbw, self.Kftrans,
576
- self.Kftrans, self.Krtrans, self.Krtrans, self.Krtrans]
582
+ extra_names, extra_counts = self._beta_segment_extra()
583
+ beta_segment_names = (["Bf", "Br_b"] + list(extra_names)
584
+ + ["chol", "Br_w", "Bftrans",
585
+ "flmbda", "Brtrans_b", "Brtrans_w", "rlmda"])
586
+ iterations = ([self.Kf, self.Kr] + list(extra_counts)
587
+ + [self.Kchol, self.Kbw, self.Kftrans,
588
+ self.Kftrans, self.Krtrans, self.Krtrans, self.Krtrans])
577
589
  self.var_list = self.split_betas(jax_result['x'], iterations, beta_segment_names)
578
590
  self.chol_mat = self.construct_chol_mat(
579
591
  self.var_list['chol'], self.var_list['Br_w'], self.var_list['Brtrans_w'])
@@ -603,64 +615,11 @@ class MixedLogit(DiscreteChoiceModel):
603
615
  args = (self.X, self.y, self.panel_info, draws, drawstrans, self.weights, self.avail, self.batch_size)
604
616
  bounds = bnds if self.method == "L-BFGS-B" else None
605
617
  options = {'gtol': self.gtol, 'maxiter': self.maxiter, 'disp': False}
606
-
607
- # ── SD penalty wrapper: regularise SD params away from exactly zero ──
608
- # When an SD parameter hits zero, the gradient vanishes and BFGS gets stuck.
609
- # A small L2 penalty that activates below a threshold keeps SDs bounded
610
- # away from zero without affecting the well-identified estimates.
611
- if self.sd_penalty > 0 and self.Kbw > 0:
612
- # SD param indices: end of beta vector, before Kftrans and Krtrans
613
- sd_start = self.Kf + self.Kr + self.Kchol
614
- sd_end = sd_start + self.Kbw
615
- _orig_obj = self.get_loglik_gradient
616
-
617
- def _penalised_obj(betas, *obj_args):
618
- f, g = _orig_obj(betas, *obj_args)
619
- for i in range(sd_start, sd_end):
620
- if i >= len(betas):
621
- break
622
- sd_val = betas[i]
623
- if sd_val < 0.05:
624
- p = self.sd_penalty * (sd_val - 0.05) ** 2
625
- f += p
626
- if g is not None and len(g) > i:
627
- g[i] += 2.0 * self.sd_penalty * (sd_val - 0.05)
628
- return f, g
629
-
630
- obj_fn = _penalised_obj
631
- else:
632
- obj_fn = self.get_loglik_gradient
633
-
634
- result = minimise_func(obj_fn, betas, jac=self.jac, method=self.method,
618
+ result = minimise_func(self.get_loglik_gradient, betas, jac=self.jac, method=self.method,
635
619
  args=args, tol=self.ftol, bounds=bounds, options=options)
636
- print(f"[MXL] Primary minimization ({self.method}): success={result.get('success', None)}, fun={result.get('fun', float('nan')):.6g}, nit={result.get('nit', '?')}")
620
+ print(f"[MXL] Minimization completed: success={result.get('success', None)}, fun={result.get('fun', float('nan')):.6g}, nit={result.get('nit', '?')}")
637
621
  if 'x' in result:
638
- print(f"[MXL] Primary betas first_values={np.asarray(result['x'])[:min(8, len(result['x']))]!r}")
639
-
640
- # ── BFGS polish: after non-BFGS methods (SLSQP, Powell, etc.),
641
- # run a quick BFGS to get the inverse Hessian for standard errors.
642
- if (self.method not in ('bfgs', 'l-bfgs-b')
643
- and 'x' in result
644
- and result.get('success', False)):
645
- try:
646
- polish_opts = {'gtol': self.gtol * 10, 'maxiter': min(self.maxiter // 5, 200), 'disp': False}
647
- polish_bnds = bnds # re-use same bounds
648
- # Use original objective (no SD penalty) for clean Hessian
649
- polish_result = minimise_func(
650
- self.get_loglik_gradient,
651
- np.asarray(result['x']).copy(),
652
- jac=True, method='L-BFGS-B',
653
- args=args, tol=self.ftol,
654
- bounds=polish_bnds, options=polish_opts,
655
- )
656
- if polish_result.get('fun', float('inf')) < result.get('fun', float('inf')) + 0.1:
657
- print(f"[MXL] BFGS polish: success={polish_result.get('success')}, "
658
- f"fun={polish_result.get('fun', float('nan')):.6g}")
659
- result = polish_result
660
- else:
661
- print(f"[MXL] BFGS polish skipped (worse objective)")
662
- except Exception as e:
663
- print(f"[MXL] BFGS polish failed: {e}")
622
+ print(f"[MXL] Minimization final betas first_values={np.asarray(result['x'])[:min(8, len(result['x']))]!r}")
664
623
  # ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
665
624
 
666
625
  if hasattr(self, 'method') and self.method == "L-BFGS-B": # {
@@ -720,7 +679,7 @@ class MixedLogit(DiscreteChoiceModel):
720
679
  @staticmethod
721
680
  def _jax_mxl_negloglik(betas, X_jax, y_jax, panel_info_jax, draws_jax,
722
681
  fxidx, rvidx, Kf, Kr, Kchol, Kbw, rvdist_names,
723
- correlationLength):
682
+ correlationLength, **kwargs):
724
683
  """JAX simulation-based log-likelihood for Mixed Logit (standard case).
725
684
 
726
685
  Handles fixed and random (normally/lognormally distributed) parameters.
@@ -849,12 +808,14 @@ class MixedLogit(DiscreteChoiceModel):
849
808
  _N = X_jax.shape[0]
850
809
  _P = X_jax.shape[1] if X_jax.ndim > 1 else 1
851
810
  _J = X_jax.shape[2] if X_jax.ndim > 2 else 1
852
- _cache_key = (_N, _P, _J, Kf, Kr, Kchol, Kbw)
811
+ _cache_key = (_N, _P, _J, Kf, Kr, Kchol, Kbw) + self._jax_cache_key_extra()
853
812
  _compiled = self._mxl_jit_cache.get(_cache_key)
854
813
  if _compiled is None:
814
+ _extra = self._jax_negloglik_extra_kwargs()
855
815
  _fn = lambda b, _X, _y, _pi, _dr: self._jax_mxl_negloglik(
856
816
  b, _X, _y, _pi, _dr,
857
- fxidx, rvidx, Kf, Kr, Kchol, Kbw, rvdist_names, correlationLength)
817
+ fxidx, rvidx, Kf, Kr, Kchol, Kbw, rvdist_names, correlationLength,
818
+ **_extra)
858
819
  _compiled = jax.jit(jax.value_and_grad(_fn))
859
820
  self._mxl_jit_cache[_cache_key] = _compiled
860
821
  # ────────────────────────────────────────────────────────
@@ -1668,4 +1629,33 @@ class MixedLogit(DiscreteChoiceModel):
1668
1629
  pch2_sd_test[i, :] = np.std(pch2[batch_slice, :], axis=0)
1669
1630
  # }
1670
1631
  self.pch2_res, self.pch2_sd_test = pch2_res, pch2_sd_test
1671
- # }
1632
+ # }
1633
+
1634
+ # ── Extension hooks for subclasses (e.g. MixedLogitGSE) ────────
1635
+ def _n_coeff_extra(self) -> int:
1636
+ """Extra parameter count beyond base MXL params."""
1637
+ return 0
1638
+
1639
+ def _bound_extra(self) -> dict:
1640
+ """Extra bound dictionary entries for subclasses."""
1641
+ return {}
1642
+
1643
+ def _beta_segment_extra(self):
1644
+ """Returns (extra_names, extra_counts) for split_betas."""
1645
+ return [], []
1646
+
1647
+ def _augment_br(self, Br, var_list):
1648
+ """Modify random coefficients. Subclasses may add gradient terms."""
1649
+ return Br
1650
+
1651
+ def _jax_negloglik_extra_kwargs(self):
1652
+ """Extra keyword arguments forwarded to JAX negloglik fn."""
1653
+ return {}
1654
+
1655
+ def _jax_cache_key_extra(self):
1656
+ """Extra tuple elements for the JAX compilation cache key."""
1657
+ return ()
1658
+
1659
+ def _init_pad_arrays(self):
1660
+ """Extra arrays to insert between Br_b and Bftrans in init_coeff."""
1661
+ return []
@@ -0,0 +1,221 @@
1
+ """MixedLogitGSE — Gradient Score Enhanced Mixed Logit.
2
+
3
+ Extends MixedLogit with gradient-latent loadings.
4
+
5
+ beta_nk = mu_k + (gamma_k + tau_k * z_nk) * g_nk + sigma_k * eta_nk
6
+
7
+ Parameter layout (no transforms, no correlation):
8
+ [Br_b(Kr) | gamma(Kr) | gamma_w(Kr) | Br_w(Kr)] (12 or 16 params)
9
+
10
+ Usage:
11
+ mxl = MixedLogitGSE()
12
+ mxl.setup(..., gradient_scores=g_avg, random_gamma=False)
13
+ mxl.fit()
14
+ """
15
+
16
+ import numpy as np
17
+ import jax.numpy as jnp
18
+ import jax
19
+
20
+ from SearchLibrium.MixedLogit import MixedLogit
21
+
22
+ try:
23
+ from ._device import device as dev
24
+ except ImportError:
25
+ from _device import device as dev
26
+
27
+
28
+ class MixedLogitGSE(MixedLogit):
29
+ """Mixed Logit with Gradient Score Enhanced latent loadings."""
30
+
31
+ def __init__(self, *args, **kwargs):
32
+ super().__init__(*args, **kwargs)
33
+ self.Kgrad = 0
34
+ self.Kgrad_w = 0
35
+ self.gamma_draws = None
36
+
37
+ def setup(self, *args, gradient_scores=None, random_gamma=False, **kwargs):
38
+ self.random_gamma = bool(random_gamma)
39
+ if gradient_scores is not None:
40
+ # Standardise to O(1) to prevent numerical overflow
41
+ g = np.asarray(gradient_scores, dtype=float)
42
+ self.gradient_scores = g / (g.std(axis=0, keepdims=True) + 1e-8)
43
+ gstd = g.std(axis=0)
44
+ print(f"[MixedLogitGSE] Gradient loadings enabled (Kgrad={len(kwargs.get('randvars', {}))}, "
45
+ f"random_gamma={self.random_gamma}, "
46
+ f"gradient std range=[{gstd.min():.2f}, {gstd.max():.2f}] -> standardised)")
47
+ else:
48
+ self.gradient_scores = None
49
+ super().setup(*args, **kwargs)
50
+
51
+ # ── Hook overrides ─────────────────────────────────────────────
52
+
53
+ def _n_coeff_extra(self) -> int:
54
+ return self.Kgrad + self.Kgrad_w
55
+
56
+ def _bound_extra(self) -> dict:
57
+ inf = float("inf")
58
+ return {"grad": ((-inf, inf), self.Kgrad),
59
+ "grad_w": ((0, inf), self.Kgrad_w)}
60
+
61
+ def _beta_segment_extra(self):
62
+ return (["grad", "grad_w"], [self.Kgrad, self.Kgrad_w])
63
+
64
+ def _jax_cache_key_extra(self):
65
+ return (self.Kgrad, self.Kgrad_w)
66
+
67
+ def _init_pad_arrays(self):
68
+ pads = []
69
+ if self.Kgrad > 0:
70
+ pads.append(np.repeat(0.1, self.Kgrad))
71
+ if self.Kgrad_w > 0:
72
+ pads.append(np.repeat(0.1, self.Kgrad_w))
73
+ return pads
74
+
75
+ def _jax_negloglik_extra_kwargs(self):
76
+ gd = (jnp.array(self.gamma_draws) if self.Kgrad_w > 0
77
+ and self.gamma_draws is not None else None)
78
+ gs = (jnp.array(self.gradient_scores) if self.gradient_scores is not None
79
+ else None)
80
+ return dict(Kgrad=self.Kgrad, Kgrad_w=self.Kgrad_w,
81
+ gradient_scores=gs, gamma_draws=gd)
82
+
83
+ # ── Fit override ───────────────────────────────────────────────
84
+
85
+ def fit(self):
86
+ if self.gradient_scores is not None:
87
+ self.Kgrad = self.Kr
88
+ self.Kgrad_w = self.Kr if self.random_gamma else 0
89
+ if self.Kgrad_w > 0:
90
+ self.gamma_draws = np.random.randn(self.N, self.Kgrad)
91
+ # Force JAX-only (scipy fallback doesn't handle gradient loadings)
92
+ self._jax = True
93
+ super().fit()
94
+
95
+ # ── JAX loglik with GSE terms ──────────────────────────────────
96
+
97
+ @staticmethod
98
+ def _jax_mxl_negloglik(betas, X_jax, y_jax, panel_info_jax, draws_jax,
99
+ fxidx, rvidx, Kf, Kr, Kchol, Kbw, rvdist_names,
100
+ correlationLength,
101
+ Kgrad=0, Kgrad_w=0,
102
+ gradient_scores=None, gamma_draws=None,
103
+ **kwargs):
104
+ """GSE loglik: β_nk = μ_k + γ_k·g_nk + σ_k·η_nk."""
105
+ # Split
106
+ Bf = betas[:Kf]
107
+ Br_b = betas[Kf:Kf + Kr]
108
+ gamma_mu = betas[Kf + Kr:Kf + Kr + Kgrad] if Kgrad > 0 else jnp.array([])
109
+ gamma_w = (betas[Kf + Kr + Kgrad:Kf + Kr + Kgrad + Kgrad_w]
110
+ if Kgrad_w > 0 else jnp.array([]))
111
+ offset = Kf + Kr + Kgrad + Kgrad_w
112
+ chol_v = betas[offset:offset + Kchol]
113
+ Br_w = betas[offset + Kchol:offset + Kchol + Kbw]
114
+
115
+ # Cholesky
116
+ chol_mat = jnp.zeros((Kr, Kr))
117
+ idx = 0
118
+ for r in range(correlationLength):
119
+ for c in range(r + 1):
120
+ chol_mat = chol_mat.at[r, c].set(chol_v[idx]); idx += 1
121
+ for k in range(Kbw):
122
+ chol_mat = chol_mat.at[correlationLength + k,
123
+ correlationLength + k].set(jnp.abs(Br_w[k]))
124
+
125
+ # Random coefficients
126
+ N = X_jax.shape[0]
127
+ Br = Br_b[:, None] + jnp.einsum("kl,nlr->nkr", chol_mat, draws_jax[:, :Kr, :])
128
+
129
+ # GSE gradient-latent term (standardised scores, O(1))
130
+ if Kgrad > 0 and gradient_scores is not None:
131
+ if Kgrad_w > 0 and gamma_draws is not None:
132
+ gamma_rnd = gamma_mu[None, :] + gamma_w[None, :] * gamma_draws
133
+ else:
134
+ gamma_rnd = gamma_mu[None, :]
135
+ Br = Br + gamma_rnd[:, :, None] * gradient_scores[:, :, None]
136
+
137
+ # Distribution transforms
138
+ for k, dist in enumerate(rvdist_names):
139
+ if dist == "ln":
140
+ Br = Br.at[:, k, :].set(jnp.exp(Br[:, k, :]))
141
+ elif dist == "tn":
142
+ Br = Br.at[:, k, :].set(jnp.abs(Br[:, k, :]))
143
+ elif dist == "u":
144
+ Br = Br.at[:, k, :].set(Br_b[k] + Br_w[k] * (draws_jax[:, k, :] - 0.5))
145
+
146
+ # Utility
147
+ Xr = X_jax[:, :, :, rvidx]
148
+ if Kf > 0:
149
+ Xf = X_jax[:, :, :, fxidx]
150
+ V = (jnp.einsum("npjk,k->npj", Xf, Bf)[:, :, :, None]
151
+ + jnp.einsum("npjk,nkr->npjr", Xr, Br))
152
+ else:
153
+ V = jnp.einsum("npjk,nkr->npjr", Xr, Br)
154
+
155
+ V = V - jnp.max(V, axis=2, keepdims=True)
156
+ eV = jnp.exp(V)
157
+ p = eV / jnp.sum(eV, axis=2, keepdims=True)
158
+
159
+ pch = jnp.sum(y_jax[:, :, :, None] * p, axis=2)
160
+ pch = jnp.prod(pch, axis=1)
161
+ pch = jnp.clip(pch, 1e-300, None)
162
+ sim_p = jnp.mean(pch, axis=1)
163
+ sim_p = jnp.clip(sim_p, 1e-300, None)
164
+ return -jnp.sum(jnp.log(sim_p))
165
+
166
+ # ── Scipy compute_probabilities override ────────────────────────
167
+
168
+ def compute_probabilities(self, betas, X, panel_info, draws, drawstrans,
169
+ avail, var_list, chol_mat):
170
+ """GSE override: add gradient-latent term to Br."""
171
+ vals = list(var_list.values())
172
+ Bf = vals[0]
173
+ Br_b = vals[1]
174
+ # GSE segments sit between Br_b and chol/Br_w
175
+ has_gse = len(vals) > 9
176
+ gamma = vals[2] if has_gse else np.array([])
177
+ gamma_w = vals[3] if has_gse else np.array([])
178
+ chol_v = vals[4] if has_gse else vals[2]
179
+ Br_w = vals[5] if has_gse else vals[3]
180
+ Bftrans = vals[6] if has_gse else vals[4] if len(vals) > 4 else np.array([])
181
+ flmbda = vals[7] if has_gse else vals[5] if len(vals) > 5 else np.array([])
182
+
183
+ if dev.using_gpu:
184
+ Bf = dev.convert_array_gpu(Bf)
185
+ Br_b = dev.convert_array_gpu(Br_b)
186
+ Br_w = dev.convert_array_gpu(Br_w)
187
+ if len(Bftrans) > 0: Bftrans = dev.convert_array_gpu(Bftrans)
188
+ if len(flmbda) > 0: flmbda = dev.convert_array_gpu(flmbda)
189
+ if len(gamma) > 0:
190
+ gamma = dev.convert_array_gpu(gamma)
191
+ gamma_w = dev.convert_array_gpu(gamma_w)
192
+
193
+ XBf = np.zeros((self.N, self.P, self.J))
194
+ if dev.using_gpu: XBf = dev.convert_array_gpu(XBf)
195
+
196
+ if self.Kf != 0:
197
+ Xf = X[:, :, :, self.fxidx]
198
+ XBf = dev.cust_einsum("npjk,k -> npj", Xf, Bf).astype(float)
199
+
200
+ XBr = np.zeros((self.N, self.P, self.J, draws.shape[2]))
201
+ if self.Kr != 0:
202
+ tmp = dev.np.matmul(chol_mat[:self.Kr, :self.Kr], draws)
203
+ Br = Br_b[None, :, None] + tmp
204
+
205
+ if self.Kgrad > 0 and self.gradient_scores is not None:
206
+ if (self.Kgrad_w > 0 and self.gamma_draws is not None
207
+ and len(gamma_w) > 0):
208
+ g_rnd = gamma[None, :] + gamma_w[None, :] * self.gamma_draws
209
+ else:
210
+ g_rnd = gamma[None, :] if len(gamma) > 0 else np.zeros((1, self.Kr))
211
+ Br = Br + g_rnd[:, :, None] * self.gradient_scores[:, :, None]
212
+
213
+ Br = self.draws_generator.apply_distribution(Br, self.rvdist)
214
+ self.Br = Br
215
+ Xr = X[:, :, :, self.rvidx].astype(float)
216
+ XBr = dev.cust_einsum("npjk,nkr -> npjr", Xr, Br)
217
+
218
+ V = XBf[:, :, :, None] + XBr if self.Kf != 0 else XBr
219
+ eV = np.exp(np.clip(V, -700, 700))
220
+ probs = eV / np.clip(eV.sum(axis=2, keepdims=True), 1e-300, None)
221
+ return probs
@@ -0,0 +1 @@
1
+ 0.0.153
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: SearchLibrium
3
- Version: 0.0.152
3
+ Version: 0.0.153
4
4
  Summary: A Python package for econometric models driven by search
5
5
  Author: Mario Fernando Taco Morales, Alexander Paz, Prithvi Beeramole, Robert Burdett
6
6
  Author-email: Zeke Ahern <z.ahern@qut.edu.au>
@@ -4,6 +4,7 @@ pyproject.toml
4
4
  src/SearchLibrium/BEHier.py
5
5
  src/SearchLibrium/Halton.py
6
6
  src/SearchLibrium/MixedLogit.py
7
+ src/SearchLibrium/MixedLogitGSE.py
7
8
  src/SearchLibrium/Mode_Activity_Nested.py
8
9
  src/SearchLibrium/RandomP.py
9
10
  src/SearchLibrium/SEARCH_SM_MARIO.py
@@ -54,4 +55,7 @@ src/SearchLibrium/bhhh/minimize.py
54
55
  src/SearchLibrium/data/Swissmetro_final.csv
55
56
  src/SearchLibrium/data/TravelMode.csv
56
57
  src/SearchLibrium/data/electricity.csv
57
- src/data_generator/latent_class_gen.py
58
+ src/data_generator/demonstrate_generation.py
59
+ src/data_generator/fit_recovery.py
60
+ src/data_generator/latent_class_gen.py
61
+ src/data_generator/print_synthetic_data.py
@@ -0,0 +1,85 @@
1
+ """
2
+ demonstrate_generation.py
3
+ =========================
4
+ Demonstrates how to generate latent class choice data with membership variables.
5
+ Run this to inspect the generated dataset and parameter structure.
6
+ """
7
+
8
+ import sys
9
+ import os
10
+ import numpy as np
11
+
12
+ sys.path.insert(0, os.path.join(os.path.dirname(__file__), ".."))
13
+ from data_generator.latent_class_gen import AdvancedLatentClassGenerator
14
+
15
+
16
+ def main(seed=42):
17
+ print("=" * 70)
18
+ print(" Data Generation Demo (with Membership Variables)")
19
+ print("=" * 70)
20
+
21
+ gen = AdvancedLatentClassGenerator(
22
+ n_classes=3,
23
+ n_alternatives=3,
24
+ n_individuals=500,
25
+ n_choice_tasks=2,
26
+ scale_separation=2.5,
27
+ n_noise_vars=2,
28
+ n_weak_vars=1,
29
+ n_collinear_vars=1,
30
+ n_membership_vars=3,
31
+ membership_scale=2.0,
32
+ random_state=seed,
33
+ )
34
+
35
+ print(f"\n Classes: {gen.K}")
36
+ print(f" Alternatives: {gen.J}")
37
+ print(f" Individuals: {gen.N}")
38
+ print(f" Choice tasks: {gen.T}")
39
+ print(f" Shared vars: {gen.shared_vars}")
40
+ print(f" Membership vars: {gen.membership_var_names}")
41
+ print(f" Noise vars: {gen.noise_vars}")
42
+ print(f" Weak vars: {gen.weak_vars}")
43
+ print(f" Collinear vars: {gen.collinear_vars}")
44
+
45
+ print(f"\n Class-specific utility vars:")
46
+ for k in range(gen.K):
47
+ print(f" Class {k}: {gen.class_specific_vars.get(k, [])}")
48
+
49
+ print(f"\n True Utility Coefficients (betas):")
50
+ for k in range(gen.K):
51
+ print(f" Class {k}:")
52
+ for v, b in gen.parameters[k].items():
53
+ print(f" {v:>25s} = {b:+.4f}")
54
+
55
+ print(f"\n True Membership Coefficients (gammas):")
56
+ if gen.gammas is not None:
57
+ for c in range(gen.K - 1):
58
+ for m in range(gen.n_membership_vars):
59
+ print(f" Class_{c + 1}_{gen.membership_var_names[m]} = {gen.gammas[c, m]:+.4f}")
60
+ print(f" Class_{gen.K} (reference) all = 0")
61
+
62
+ df, true_classes = gen.generate()
63
+
64
+ print(f"\n Generated data: {len(df)} rows")
65
+ print(f" Columns: {list(df.columns)}")
66
+ print(f" Class distribution:")
67
+ for k in range(gen.K):
68
+ n_k = (true_classes == k).sum()
69
+ print(f" Class {k}: {n_k} ({n_k / gen.N:.1%})")
70
+
71
+ print(f"\n First 5 rows:")
72
+ print(df.head(10).to_string())
73
+
74
+ print(f"\n Individual-level membership vars (first 5 individuals):")
75
+ mem_cols = gen.membership_var_names
76
+ for n in range(min(5, gen.N)):
77
+ vals = {v: df.loc[df["individual"] == n, v].iloc[0] for v in mem_cols}
78
+ print(f" Individual {n} (true_class={df.loc[df['individual'] == n, 'true_class'].iloc[0]}): {vals}")
79
+
80
+ print(f"\n Done. Use the generated data with SearchLibrium's LatentClassMixedLogit.")
81
+ return gen, df, true_classes
82
+
83
+
84
+ if __name__ == "__main__":
85
+ main()