Repository files navigation

Jaxlogit

MSLE estimation of linear-in-parameters mixed logit models. This package is based on xlogit, with the core computational engine replaced by jax. Additionally, the documentation structure and many examples are adapted from xlogit's documentation.

Why use Jaxlogit?

While there are many packages out there that can compute mixed logit models, jaxlogit is unique in its processing speed and capabilities, which come from using jax and just in time compilation. You can see the comparison here. Because jaxlogit specialises in linear models, it can use accelerated linear algebra to do the matrix multiplication, which significantly increases the speed. Jaxlogit can also work with large quanities of data by calculating them in batches, reducing memory load. This can be seen in this example.

Quick Start

This example uses jaxlogit to to estimate a mixed logit model for choices of transport method using the swissmetro dataset. It contains stated-preferences for three alternative transportation modes that include car, train and a newly introduced mode: the swissmetro. The dataset is available here and Bierlaire et. al., (2001) provides a detailed discussion of the data as well as its context and collection process.

The explanatory variables are cost and travel time of each of the alternatives, as well as alternative-specific constraints to represent unobserved factors. New functionality used compared to xlogit is considering the correlation between the distributions of explanatory variables. In this example, only correlations between some variables are considered, and some are set to certain values, according to research done in J. Walker's PhD thesis (MIT 2001). This is specified by specifying some correlation parameters to be 0.

The data must be in long format. Wide format data can be transformed into long format data using the wide_to_long function.

Required parameters:

  • X: 2-D array of input data (in long format) with choice situations as rows, and variables as columns
  • y: 1-D array of choices (in long format)
  • varnames: List of variable names that matches the number and order of the columns in X
  • alts: 1-D array of alternative representing the alternative chosen for each input
  • ids: 1-D array of the ids of the choice situations
  • randvars: dictionary of variables and their mixing distributions. In this example, only "n" normal variables are used. Valid distributions are "n" normal, "ln" log normal and "n_trunc" truncated normal.

Optional Parameters in this example:

  • avail: 2-D array of availability of alternatives for the choice situations. One when available or zero otherwise.
  • panels: 1-D array of ids for panel formation, where each choice situation has a panel for each variable
  • set_vars: Dictionary specifying some variables and the values that they are to be set
  • include_correlations: Boolean whether or not to consider correlation between variables
  • optim_method: String representing with optimisation method to use. Options are "L-BFGS-scipy", "BFGS-scipy", "L-BFGS-jax", and "BFGS-jax". Scipy uses the standard scipy library and jax uses jax's scipy library, which is significantly faster (may not work with batching) but may not be maintained and potentially discontinued without notice.
importpandasaspdimportnumpyasnpimportjaxfromjaxlogit.mixed_logitimportMixedLogit, ConfigData# 64bit precisionjax.config.update("jax_enable_x64", True)
# read and format datadf_wide=pd.read_table("http://transp-or.epfl.ch/data/swissmetro.dat", sep='\t')
# Keep only observations for commute and business purposes that contain known choicesdf_wide=df_wide[(df_wide['PURPOSE'].isin([1, 3]) & (df_wide['CHOICE'] !=0))]
df_wide['custom_id'] =np.arange(len(df_wide)) # Add unique identifierdf_wide['CHOICE'] =df_wide['CHOICE'].map({1: 'TRAIN', 2:'SM', 3: 'CAR'})
fromjaxlogit.utilsimportwide_to_longdf=wide_to_long(df_wide, id_col='custom_id', alt_name='alt', sep='_',
alt_list=['TRAIN', 'SM', 'CAR'], empty_val=0,
varying=['TT', 'CO', 'HE', 'AV', 'SEATS'], alt_is_prefix=True)
# modification of data to includedf['ASC_TRAIN'] =np.ones(len(df))*(df['alt'] =='TRAIN')
df['ASC_CAR'] =np.ones(len(df))*(df['alt'] =='CAR')
df['TT'], df['CO'] =df['TT']/100, df['CO']/100# Scale variablesannual_pass= (df['GA'] ==1) & (df['alt'].isin(['TRAIN', 'SM']))
df.loc[annual_pass, 'CO'] =0# Cost zero for pass holders# specification of variablesvarnames=['ASC_CAR', 'ASC_TRAIN', 'ASC_SM', 'CO', 'TT']
randvars={'ASC_CAR': 'n', 'ASC_TRAIN': 'n', 'ASC_SM': 'n'}
set_vars= {
'ASC_SM': 0.0, 'sd.ASC_TRAIN': 1.0, 'sd.ASC_CAR': 0.0,
'chol.ASC_CAR.ASC_TRAIN': 0.0, 'chol.ASC_CAR.ASC_SM': 0.0
} # Identification of error components, see J. Walker's PhD thesis (MIT 2001)# change some optional argumentsconfig=ConfigData(
avail=df['AV'],
panels=df["ID"],
set_vars=set_vars,
include_correlations=True, # Enable correlation between random parametersoptim_method="L-BFGS-scipy"
)
model=MixedLogit()
res=model.fit(
X=df[varnames],
y=df['CHOICE'],
varnames=varnames,
alts=df['alt'],
ids=df['custom_id'],
randvars=randvars,
config=config
)
model.summary()
 Message: CONVERGENCE: RELATIVE REDUCTION OF F <= FACTR*EPSMCH
Iterations: 17
Function evaluations: 21
Estimation time= 57.3 seconds
---------------------------------------------------------------------------
Coefficient Estimate Std.Err. z-val P>|z|
---------------------------------------------------------------------------
ASC_CAR -0.2279764 0.4655508 -0.4896917 0.624 ASC_TRAIN -1.1966263 0.5873328 -2.0373905 0.0416 * ASC_SM 0.1000000 0.0000000 inf 0 ***
CO -2.0490264 0.3444999 -5.9478289 2.85e-09 ***
TT -2.1429727 0.6607183 -3.2433986 0.00119 ** sd.ASC_CAR 0.1000000 0.0000000 inf 0 ***
sd.ASC_TRAIN 0.1000000 0.0000000 inf 0 ***
sd.ASC_SM 2.4459162 0.2814464 8.6905216 4.47e-18 ***
chol.ASC_CAR.ASC_TR 0.1000000 0.0000000 inf 0 ***
chol.ASC_CAR.ASC_SM 0.1000000 0.0000000 inf 0 ***
chol.ASC_TRAIN.ASC_ 0.6937107 0.2260149 3.0693135 0.00215 ** ---------------------------------------------------------------------------
Significance: 0 '***' 0.001 '**' 0.01 '*' 0.05 '.' 0.1 ' ' 1
Log-Likelihood= -4071.438
AIC= 8154.876
BIC= 8195.795

If out of memory, the data can be batched as well.

Quick Install

Install jaxlogit using pip: pip install jaxlogit.

Alternatively, clone the repo.

Benchmark

As shown in the plot below, jaxlogit with batching uses significantly less memory than xlogit. Graph comparing memory usage of jaxlogit and xlogit

Graph comparing time of jaxlogit, xlogit, and biogeme

The graph of memory shows that using the batching reduces memory usage below that of xlogit and other types of jaxlogit. Off the edge of this graph is a spike in biogeme's memory usage, up to 30GB.

In the timing graph, it can be seen that batching does not have a significant impact on time taken. Additionally, the jaxlogit performs much faster when using the experimental jax methods. The time for standard scipy method and xlogit are quite comparable. Biogeme is slower than xlogit and any form of jaxlogit.

These were taken on the complicated electricity dataset. On the nicer artificial dataset the timing looks like this: Graph comparing timing of jaxlogit and xlogit

About

No description, website, or topics provided.

Resources

Stars

2 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages

, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Add copy buttons to all
 blocks\n(function() {\n function addCopyButtons() {\n document.querySelectorAll('pre code').forEach(function(codeBlock) {\n if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;\n codeBlock.parentElement.setAttribute('data-copy-added', 'true');\n \n var btn = document.createElement('button');\n btn.textContent = 'Copy';\n btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';\n btn.onmouseover = function() { this.style.opacity = '1'; };\n btn.onmouseout = function() { this.style.opacity = '0.7'; };\n btn.onclick = function() {\n navigator.clipboard.writeText(codeBlock.textContent).then(function() {\n btn.textContent = 'Copied!';\n setTimeout(function() { btn.textContent = 'Copy'; }, 1500);\n });\n };\n codeBlock.parentElement.style.position = 'relative';\n codeBlock.parentElement.appendChild(btn);\n });\n }\n \n addCopyButtons();\n \n // Re-run on dynamic content\n var observer = new MutationObserver(addCopyButtons);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Add Copy Buttons to Code Blocks");
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
Skip to content

Repository files navigation

Jaxlogit

MSLE estimation of linear-in-parameters mixed logit models. This package is based on xlogit, with the core computational engine replaced by jax. Additionally, the documentation structure and many examples are adapted from xlogit's documentation.

Why use Jaxlogit?

While there are many packages out there that can compute mixed logit models, jaxlogit is unique in its processing speed and capabilities, which come from using jax and just in time compilation. You can see the comparison here. Because jaxlogit specialises in linear models, it can use accelerated linear algebra to do the matrix multiplication, which significantly increases the speed. Jaxlogit can also work with large quanities of data by calculating them in batches, reducing memory load. This can be seen in this example.

Quick Start

This example uses jaxlogit to to estimate a mixed logit model for choices of transport method using the swissmetro dataset. It contains stated-preferences for three alternative transportation modes that include car, train and a newly introduced mode: the swissmetro. The dataset is available here and Bierlaire et. al., (2001) provides a detailed discussion of the data as well as its context and collection process.

The explanatory variables are cost and travel time of each of the alternatives, as well as alternative-specific constraints to represent unobserved factors. New functionality used compared to xlogit is considering the correlation between the distributions of explanatory variables. In this example, only correlations between some variables are considered, and some are set to certain values, according to research done in J. Walker's PhD thesis (MIT 2001). This is specified by specifying some correlation parameters to be 0.

The data must be in long format. Wide format data can be transformed into long format data using the wide_to_long function.

Required parameters:

  • X: 2-D array of input data (in long format) with choice situations as rows, and variables as columns
  • y: 1-D array of choices (in long format)
  • varnames: List of variable names that matches the number and order of the columns in X
  • alts: 1-D array of alternative representing the alternative chosen for each input
  • ids: 1-D array of the ids of the choice situations
  • randvars: dictionary of variables and their mixing distributions. In this example, only "n" normal variables are used. Valid distributions are "n" normal, "ln" log normal and "n_trunc" truncated normal.

Optional Parameters in this example:

  • avail: 2-D array of availability of alternatives for the choice situations. One when available or zero otherwise.
  • panels: 1-D array of ids for panel formation, where each choice situation has a panel for each variable
  • set_vars: Dictionary specifying some variables and the values that they are to be set
  • include_correlations: Boolean whether or not to consider correlation between variables
  • optim_method: String representing with optimisation method to use. Options are "L-BFGS-scipy", "BFGS-scipy", "L-BFGS-jax", and "BFGS-jax". Scipy uses the standard scipy library and jax uses jax's scipy library, which is significantly faster (may not work with batching) but may not be maintained and potentially discontinued without notice.
importpandasaspdimportnumpyasnpimportjaxfromjaxlogit.mixed_logitimportMixedLogit, ConfigData# 64bit precisionjax.config.update("jax_enable_x64", True)
# read and format datadf_wide=pd.read_table("http://transp-or.epfl.ch/data/swissmetro.dat", sep='\t')
# Keep only observations for commute and business purposes that contain known choicesdf_wide=df_wide[(df_wide['PURPOSE'].isin([1, 3]) & (df_wide['CHOICE'] !=0))]
df_wide['custom_id'] =np.arange(len(df_wide)) # Add unique identifierdf_wide['CHOICE'] =df_wide['CHOICE'].map({1: 'TRAIN', 2:'SM', 3: 'CAR'})
fromjaxlogit.utilsimportwide_to_longdf=wide_to_long(df_wide, id_col='custom_id', alt_name='alt', sep='_',
alt_list=['TRAIN', 'SM', 'CAR'], empty_val=0,
varying=['TT', 'CO', 'HE', 'AV', 'SEATS'], alt_is_prefix=True)
# modification of data to includedf['ASC_TRAIN'] =np.ones(len(df))*(df['alt'] =='TRAIN')
df['ASC_CAR'] =np.ones(len(df))*(df['alt'] =='CAR')
df['TT'], df['CO'] =df['TT']/100, df['CO']/100# Scale variablesannual_pass= (df['GA'] ==1) & (df['alt'].isin(['TRAIN', 'SM']))
df.loc[annual_pass, 'CO'] =0# Cost zero for pass holders# specification of variablesvarnames=['ASC_CAR', 'ASC_TRAIN', 'ASC_SM', 'CO', 'TT']
randvars={'ASC_CAR': 'n', 'ASC_TRAIN': 'n', 'ASC_SM': 'n'}
set_vars= {
'ASC_SM': 0.0, 'sd.ASC_TRAIN': 1.0, 'sd.ASC_CAR': 0.0,
'chol.ASC_CAR.ASC_TRAIN': 0.0, 'chol.ASC_CAR.ASC_SM': 0.0
} # Identification of error components, see J. Walker's PhD thesis (MIT 2001)# change some optional argumentsconfig=ConfigData(
avail=df['AV'],
panels=df["ID"],
set_vars=set_vars,
include_correlations=True, # Enable correlation between random parametersoptim_method="L-BFGS-scipy"
)
model=MixedLogit()
res=model.fit(
X=df[varnames],
y=df['CHOICE'],
varnames=varnames,
alts=df['alt'],
ids=df['custom_id'],
randvars=randvars,
config=config
)
model.summary()
 Message: CONVERGENCE: RELATIVE REDUCTION OF F <= FACTR*EPSMCH
Iterations: 17
Function evaluations: 21
Estimation time= 57.3 seconds
---------------------------------------------------------------------------
Coefficient Estimate Std.Err. z-val P>|z|
---------------------------------------------------------------------------
ASC_CAR -0.2279764 0.4655508 -0.4896917 0.624 ASC_TRAIN -1.1966263 0.5873328 -2.0373905 0.0416 * ASC_SM 0.1000000 0.0000000 inf 0 ***
CO -2.0490264 0.3444999 -5.9478289 2.85e-09 ***
TT -2.1429727 0.6607183 -3.2433986 0.00119 ** sd.ASC_CAR 0.1000000 0.0000000 inf 0 ***
sd.ASC_TRAIN 0.1000000 0.0000000 inf 0 ***
sd.ASC_SM 2.4459162 0.2814464 8.6905216 4.47e-18 ***
chol.ASC_CAR.ASC_TR 0.1000000 0.0000000 inf 0 ***
chol.ASC_CAR.ASC_SM 0.1000000 0.0000000 inf 0 ***
chol.ASC_TRAIN.ASC_ 0.6937107 0.2260149 3.0693135 0.00215 ** ---------------------------------------------------------------------------
Significance: 0 '***' 0.001 '**' 0.01 '*' 0.05 '.' 0.1 ' ' 1
Log-Likelihood= -4071.438
AIC= 8154.876
BIC= 8195.795

If out of memory, the data can be batched as well.

Quick Install

Install jaxlogit using pip: pip install jaxlogit.

Alternatively, clone the repo.

Benchmark

As shown in the plot below, jaxlogit with batching uses significantly less memory than xlogit. Graph comparing memory usage of jaxlogit and xlogit

Graph comparing time of jaxlogit, xlogit, and biogeme

The graph of memory shows that using the batching reduces memory usage below that of xlogit and other types of jaxlogit. Off the edge of this graph is a spike in biogeme's memory usage, up to 30GB.

In the timing graph, it can be seen that batching does not have a significant impact on time taken. Additionally, the jaxlogit performs much faster when using the experimental jax methods. The time for standard scipy method and xlogit are quite comparable. Biogeme is slower than xlogit and any form of jaxlogit.

These were taken on the complicated electricity dataset. On the nicer artificial dataset the timing looks like this: Graph comparing timing of jaxlogit and xlogit

About

No description, website, or topics provided.

Resources

Stars

2 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages

, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Force GitHub README to respect dark mode\n(function() {\n var style = document.createElement('style');\n style.textContent = '\n .markdown-body {\n color-scheme: dark light;\n }\n .markdown-body pre { background: #161b22 !important; }\n .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; }\n .markdown-body table th, .markdown-body table td { border-color: #30363d !important; }\n .markdown-body img { background: #0d1117; }\n .markdown-body blockquote { border-left-color: #8b949e; }\n .markdown-body hr { border-color: #30363d; }\n ';\n document.head.appendChild(style);\n})();", "GitHub Dark Mode README Fix"); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content

Repository files navigation

Jaxlogit

MSLE estimation of linear-in-parameters mixed logit models. This package is based on xlogit, with the core computational engine replaced by jax. Additionally, the documentation structure and many examples are adapted from xlogit's documentation.

Why use Jaxlogit?

While there are many packages out there that can compute mixed logit models, jaxlogit is unique in its processing speed and capabilities, which come from using jax and just in time compilation. You can see the comparison here. Because jaxlogit specialises in linear models, it can use accelerated linear algebra to do the matrix multiplication, which significantly increases the speed. Jaxlogit can also work with large quanities of data by calculating them in batches, reducing memory load. This can be seen in this example.

Quick Start

This example uses jaxlogit to to estimate a mixed logit model for choices of transport method using the swissmetro dataset. It contains stated-preferences for three alternative transportation modes that include car, train and a newly introduced mode: the swissmetro. The dataset is available here and Bierlaire et. al., (2001) provides a detailed discussion of the data as well as its context and collection process.

The explanatory variables are cost and travel time of each of the alternatives, as well as alternative-specific constraints to represent unobserved factors. New functionality used compared to xlogit is considering the correlation between the distributions of explanatory variables. In this example, only correlations between some variables are considered, and some are set to certain values, according to research done in J. Walker's PhD thesis (MIT 2001). This is specified by specifying some correlation parameters to be 0.

The data must be in long format. Wide format data can be transformed into long format data using the wide_to_long function.

Required parameters:

  • X: 2-D array of input data (in long format) with choice situations as rows, and variables as columns
  • y: 1-D array of choices (in long format)
  • varnames: List of variable names that matches the number and order of the columns in X
  • alts: 1-D array of alternative representing the alternative chosen for each input
  • ids: 1-D array of the ids of the choice situations
  • randvars: dictionary of variables and their mixing distributions. In this example, only "n" normal variables are used. Valid distributions are "n" normal, "ln" log normal and "n_trunc" truncated normal.

Optional Parameters in this example:

  • avail: 2-D array of availability of alternatives for the choice situations. One when available or zero otherwise.
  • panels: 1-D array of ids for panel formation, where each choice situation has a panel for each variable
  • set_vars: Dictionary specifying some variables and the values that they are to be set
  • include_correlations: Boolean whether or not to consider correlation between variables
  • optim_method: String representing with optimisation method to use. Options are "L-BFGS-scipy", "BFGS-scipy", "L-BFGS-jax", and "BFGS-jax". Scipy uses the standard scipy library and jax uses jax's scipy library, which is significantly faster (may not work with batching) but may not be maintained and potentially discontinued without notice.
importpandasaspdimportnumpyasnpimportjaxfromjaxlogit.mixed_logitimportMixedLogit, ConfigData# 64bit precisionjax.config.update("jax_enable_x64", True)
# read and format datadf_wide=pd.read_table("http://transp-or.epfl.ch/data/swissmetro.dat", sep='\t')
# Keep only observations for commute and business purposes that contain known choicesdf_wide=df_wide[(df_wide['PURPOSE'].isin([1, 3]) & (df_wide['CHOICE'] !=0))]
df_wide['custom_id'] =np.arange(len(df_wide)) # Add unique identifierdf_wide['CHOICE'] =df_wide['CHOICE'].map({1: 'TRAIN', 2:'SM', 3: 'CAR'})
fromjaxlogit.utilsimportwide_to_longdf=wide_to_long(df_wide, id_col='custom_id', alt_name='alt', sep='_',
alt_list=['TRAIN', 'SM', 'CAR'], empty_val=0,
varying=['TT', 'CO', 'HE', 'AV', 'SEATS'], alt_is_prefix=True)
# modification of data to includedf['ASC_TRAIN'] =np.ones(len(df))*(df['alt'] =='TRAIN')
df['ASC_CAR'] =np.ones(len(df))*(df['alt'] =='CAR')
df['TT'], df['CO'] =df['TT']/100, df['CO']/100# Scale variablesannual_pass= (df['GA'] ==1) & (df['alt'].isin(['TRAIN', 'SM']))
df.loc[annual_pass, 'CO'] =0# Cost zero for pass holders# specification of variablesvarnames=['ASC_CAR', 'ASC_TRAIN', 'ASC_SM', 'CO', 'TT']
randvars={'ASC_CAR': 'n', 'ASC_TRAIN': 'n', 'ASC_SM': 'n'}
set_vars= {
'ASC_SM': 0.0, 'sd.ASC_TRAIN': 1.0, 'sd.ASC_CAR': 0.0,
'chol.ASC_CAR.ASC_TRAIN': 0.0, 'chol.ASC_CAR.ASC_SM': 0.0
} # Identification of error components, see J. Walker's PhD thesis (MIT 2001)# change some optional argumentsconfig=ConfigData(
avail=df['AV'],
panels=df["ID"],
set_vars=set_vars,
include_correlations=True, # Enable correlation between random parametersoptim_method="L-BFGS-scipy"
)
model=MixedLogit()
res=model.fit(
X=df[varnames],
y=df['CHOICE'],
varnames=varnames,
alts=df['alt'],
ids=df['custom_id'],
randvars=randvars,
config=config
)
model.summary()
 Message: CONVERGENCE: RELATIVE REDUCTION OF F <= FACTR*EPSMCH
Iterations: 17
Function evaluations: 21
Estimation time= 57.3 seconds
---------------------------------------------------------------------------
Coefficient Estimate Std.Err. z-val P>|z|
---------------------------------------------------------------------------
ASC_CAR -0.2279764 0.4655508 -0.4896917 0.624 ASC_TRAIN -1.1966263 0.5873328 -2.0373905 0.0416 * ASC_SM 0.1000000 0.0000000 inf 0 ***
CO -2.0490264 0.3444999 -5.9478289 2.85e-09 ***
TT -2.1429727 0.6607183 -3.2433986 0.00119 ** sd.ASC_CAR 0.1000000 0.0000000 inf 0 ***
sd.ASC_TRAIN 0.1000000 0.0000000 inf 0 ***
sd.ASC_SM 2.4459162 0.2814464 8.6905216 4.47e-18 ***
chol.ASC_CAR.ASC_TR 0.1000000 0.0000000 inf 0 ***
chol.ASC_CAR.ASC_SM 0.1000000 0.0000000 inf 0 ***
chol.ASC_TRAIN.ASC_ 0.6937107 0.2260149 3.0693135 0.00215 ** ---------------------------------------------------------------------------
Significance: 0 '***' 0.001 '**' 0.01 '*' 0.05 '.' 0.1 ' ' 1
Log-Likelihood= -4071.438
AIC= 8154.876
BIC= 8195.795

If out of memory, the data can be batched as well.

Quick Install

Install jaxlogit using pip: pip install jaxlogit.

Alternatively, clone the repo.

Benchmark

As shown in the plot below, jaxlogit with batching uses significantly less memory than xlogit. Graph comparing memory usage of jaxlogit and xlogit

Graph comparing time of jaxlogit, xlogit, and biogeme

The graph of memory shows that using the batching reduces memory usage below that of xlogit and other types of jaxlogit. Off the edge of this graph is a spike in biogeme's memory usage, up to 30GB.

In the timing graph, it can be seen that batching does not have a significant impact on time taken. Additionally, the jaxlogit performs much faster when using the experimental jax methods. The time for standard scipy method and xlogit are quite comparable. Biogeme is slower than xlogit and any form of jaxlogit.

These were taken on the complicated electricity dataset. On the nicer artificial dataset the timing looks like this: Graph comparing timing of jaxlogit and xlogit

About

No description, website, or topics provided.

Resources

Stars

2 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages

, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Highlight search terms from Google/DuckDuckGo/Bing referrer\n(function() {\n var ref = document.referrer;\n var terms = [];\n \n if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) {\n var url = new URL(ref);\n var q = url.searchParams.get('q') || url.searchParams.get('p');\n if (q) {\n terms = q.split(/\\s+/).filter(function(t) { return t.length > 2; });\n }\n }\n \n if (terms.length === 0) return;\n \n var style = document.createElement('style');\n style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }';\n document.head.appendChild(style);\n \n function highlight(node) {\n if (node.nodeType === 3) { // text node\n var text = node.textContent;\n var found = false;\n terms.forEach(function(term) {\n var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\\]\\\\]/g, '\\\\') + ')', 'gi');\n if (regex.test(text)) {\n found = true;\n var frag = document.createDocumentFragment();\n var parts = text.split(regex);\n parts.forEach(function(part, i) {\n if (i % 2 === 0) {\n frag.appendChild(document.createTextNode(part));\n } else {\n var span = document.createElement('span');\n span.className = 'userscript-highlight';\n span.textContent = part;\n frag.appendChild(span);\n }\n });\n node.parentNode.replaceChild(frag, node);\n }\n });\n } else if (node.nodeType === 1 && node.childNodes) { // element\n var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT'];\n if (!skipTags.includes(node.tagName)) {\n Array.from(node.childNodes).forEach(highlight);\n }\n }\n }\n \n highlight(document.body);\n \n // Re-highlight on dynamic content\n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1 || node.nodeType === 3) highlight(node);\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Highlight Search Terms"); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content

Repository files navigation

Jaxlogit

MSLE estimation of linear-in-parameters mixed logit models. This package is based on xlogit, with the core computational engine replaced by jax. Additionally, the documentation structure and many examples are adapted from xlogit's documentation.

Why use Jaxlogit?

While there are many packages out there that can compute mixed logit models, jaxlogit is unique in its processing speed and capabilities, which come from using jax and just in time compilation. You can see the comparison here. Because jaxlogit specialises in linear models, it can use accelerated linear algebra to do the matrix multiplication, which significantly increases the speed. Jaxlogit can also work with large quanities of data by calculating them in batches, reducing memory load. This can be seen in this example.

Quick Start

This example uses jaxlogit to to estimate a mixed logit model for choices of transport method using the swissmetro dataset. It contains stated-preferences for three alternative transportation modes that include car, train and a newly introduced mode: the swissmetro. The dataset is available here and Bierlaire et. al., (2001) provides a detailed discussion of the data as well as its context and collection process.

The explanatory variables are cost and travel time of each of the alternatives, as well as alternative-specific constraints to represent unobserved factors. New functionality used compared to xlogit is considering the correlation between the distributions of explanatory variables. In this example, only correlations between some variables are considered, and some are set to certain values, according to research done in J. Walker's PhD thesis (MIT 2001). This is specified by specifying some correlation parameters to be 0.

The data must be in long format. Wide format data can be transformed into long format data using the wide_to_long function.

Required parameters:

  • X: 2-D array of input data (in long format) with choice situations as rows, and variables as columns
  • y: 1-D array of choices (in long format)
  • varnames: List of variable names that matches the number and order of the columns in X
  • alts: 1-D array of alternative representing the alternative chosen for each input
  • ids: 1-D array of the ids of the choice situations
  • randvars: dictionary of variables and their mixing distributions. In this example, only "n" normal variables are used. Valid distributions are "n" normal, "ln" log normal and "n_trunc" truncated normal.

Optional Parameters in this example:

  • avail: 2-D array of availability of alternatives for the choice situations. One when available or zero otherwise.
  • panels: 1-D array of ids for panel formation, where each choice situation has a panel for each variable
  • set_vars: Dictionary specifying some variables and the values that they are to be set
  • include_correlations: Boolean whether or not to consider correlation between variables
  • optim_method: String representing with optimisation method to use. Options are "L-BFGS-scipy", "BFGS-scipy", "L-BFGS-jax", and "BFGS-jax". Scipy uses the standard scipy library and jax uses jax's scipy library, which is significantly faster (may not work with batching) but may not be maintained and potentially discontinued without notice.
importpandasaspdimportnumpyasnpimportjaxfromjaxlogit.mixed_logitimportMixedLogit, ConfigData# 64bit precisionjax.config.update("jax_enable_x64", True)
# read and format datadf_wide=pd.read_table("http://transp-or.epfl.ch/data/swissmetro.dat", sep='\t')
# Keep only observations for commute and business purposes that contain known choicesdf_wide=df_wide[(df_wide['PURPOSE'].isin([1, 3]) & (df_wide['CHOICE'] !=0))]
df_wide['custom_id'] =np.arange(len(df_wide)) # Add unique identifierdf_wide['CHOICE'] =df_wide['CHOICE'].map({1: 'TRAIN', 2:'SM', 3: 'CAR'})
fromjaxlogit.utilsimportwide_to_longdf=wide_to_long(df_wide, id_col='custom_id', alt_name='alt', sep='_',
alt_list=['TRAIN', 'SM', 'CAR'], empty_val=0,
varying=['TT', 'CO', 'HE', 'AV', 'SEATS'], alt_is_prefix=True)
# modification of data to includedf['ASC_TRAIN'] =np.ones(len(df))*(df['alt'] =='TRAIN')
df['ASC_CAR'] =np.ones(len(df))*(df['alt'] =='CAR')
df['TT'], df['CO'] =df['TT']/100, df['CO']/100# Scale variablesannual_pass= (df['GA'] ==1) & (df['alt'].isin(['TRAIN', 'SM']))
df.loc[annual_pass, 'CO'] =0# Cost zero for pass holders# specification of variablesvarnames=['ASC_CAR', 'ASC_TRAIN', 'ASC_SM', 'CO', 'TT']
randvars={'ASC_CAR': 'n', 'ASC_TRAIN': 'n', 'ASC_SM': 'n'}
set_vars= {
'ASC_SM': 0.0, 'sd.ASC_TRAIN': 1.0, 'sd.ASC_CAR': 0.0,
'chol.ASC_CAR.ASC_TRAIN': 0.0, 'chol.ASC_CAR.ASC_SM': 0.0
} # Identification of error components, see J. Walker's PhD thesis (MIT 2001)# change some optional argumentsconfig=ConfigData(
avail=df['AV'],
panels=df["ID"],
set_vars=set_vars,
include_correlations=True, # Enable correlation between random parametersoptim_method="L-BFGS-scipy"
)
model=MixedLogit()
res=model.fit(
X=df[varnames],
y=df['CHOICE'],
varnames=varnames,
alts=df['alt'],
ids=df['custom_id'],
randvars=randvars,
config=config
)
model.summary()
 Message: CONVERGENCE: RELATIVE REDUCTION OF F <= FACTR*EPSMCH
Iterations: 17
Function evaluations: 21
Estimation time= 57.3 seconds
---------------------------------------------------------------------------
Coefficient Estimate Std.Err. z-val P>|z|
---------------------------------------------------------------------------
ASC_CAR -0.2279764 0.4655508 -0.4896917 0.624 ASC_TRAIN -1.1966263 0.5873328 -2.0373905 0.0416 * ASC_SM 0.1000000 0.0000000 inf 0 ***
CO -2.0490264 0.3444999 -5.9478289 2.85e-09 ***
TT -2.1429727 0.6607183 -3.2433986 0.00119 ** sd.ASC_CAR 0.1000000 0.0000000 inf 0 ***
sd.ASC_TRAIN 0.1000000 0.0000000 inf 0 ***
sd.ASC_SM 2.4459162 0.2814464 8.6905216 4.47e-18 ***
chol.ASC_CAR.ASC_TR 0.1000000 0.0000000 inf 0 ***
chol.ASC_CAR.ASC_SM 0.1000000 0.0000000 inf 0 ***
chol.ASC_TRAIN.ASC_ 0.6937107 0.2260149 3.0693135 0.00215 ** ---------------------------------------------------------------------------
Significance: 0 '***' 0.001 '**' 0.01 '*' 0.05 '.' 0.1 ' ' 1
Log-Likelihood= -4071.438
AIC= 8154.876
BIC= 8195.795

If out of memory, the data can be batched as well.

Quick Install

Install jaxlogit using pip: pip install jaxlogit.

Alternatively, clone the repo.

Benchmark

As shown in the plot below, jaxlogit with batching uses significantly less memory than xlogit. Graph comparing memory usage of jaxlogit and xlogit

Graph comparing time of jaxlogit, xlogit, and biogeme

The graph of memory shows that using the batching reduces memory usage below that of xlogit and other types of jaxlogit. Off the edge of this graph is a spike in biogeme's memory usage, up to 30GB.

In the timing graph, it can be seen that batching does not have a significant impact on time taken. Additionally, the jaxlogit performs much faster when using the experimental jax methods. The time for standard scipy method and xlogit are quite comparable. Biogeme is slower than xlogit and any form of jaxlogit.

These were taken on the complicated electricity dataset. On the nicer artificial dataset the timing looks like this: Graph comparing timing of jaxlogit and xlogit

About

No description, website, or topics provided.

Resources

Stars

2 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages

, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Strip utm_, fbclid, gclid, etc. from all links on page\n(function() {\n var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content',\n 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid',\n 'ref', 'ref_src', 'source', 'medium', 'campaign'];\n \n function cleanUrl(url) {\n try {\n var u = new URL(url, window.location.origin);\n var changed = false;\n trackingParams.forEach(function(p) {\n if (u.searchParams.has(p)) {\n u.searchParams.delete(p);\n changed = true;\n }\n });\n return changed ? u.toString() : url;\n } catch (e) {\n return url;\n }\n }\n \n function cleanLinks() {\n document.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n \n cleanLinks();\n \n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1) {\n if (node.tagName === 'A') cleanLinks();\n node.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Remove Tracking Parameters from Links"); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + '
Skip to content

Repository files navigation

Jaxlogit

MSLE estimation of linear-in-parameters mixed logit models. This package is based on xlogit, with the core computational engine replaced by jax. Additionally, the documentation structure and many examples are adapted from xlogit's documentation.

Why use Jaxlogit?

While there are many packages out there that can compute mixed logit models, jaxlogit is unique in its processing speed and capabilities, which come from using jax and just in time compilation. You can see the comparison here. Because jaxlogit specialises in linear models, it can use accelerated linear algebra to do the matrix multiplication, which significantly increases the speed. Jaxlogit can also work with large quanities of data by calculating them in batches, reducing memory load. This can be seen in this example.

Quick Start

This example uses jaxlogit to to estimate a mixed logit model for choices of transport method using the swissmetro dataset. It contains stated-preferences for three alternative transportation modes that include car, train and a newly introduced mode: the swissmetro. The dataset is available here and Bierlaire et. al., (2001) provides a detailed discussion of the data as well as its context and collection process.

The explanatory variables are cost and travel time of each of the alternatives, as well as alternative-specific constraints to represent unobserved factors. New functionality used compared to xlogit is considering the correlation between the distributions of explanatory variables. In this example, only correlations between some variables are considered, and some are set to certain values, according to research done in J. Walker's PhD thesis (MIT 2001). This is specified by specifying some correlation parameters to be 0.

The data must be in long format. Wide format data can be transformed into long format data using the wide_to_long function.

Required parameters:

  • X: 2-D array of input data (in long format) with choice situations as rows, and variables as columns
  • y: 1-D array of choices (in long format)
  • varnames: List of variable names that matches the number and order of the columns in X
  • alts: 1-D array of alternative representing the alternative chosen for each input
  • ids: 1-D array of the ids of the choice situations
  • randvars: dictionary of variables and their mixing distributions. In this example, only "n" normal variables are used. Valid distributions are "n" normal, "ln" log normal and "n_trunc" truncated normal.

Optional Parameters in this example:

  • avail: 2-D array of availability of alternatives for the choice situations. One when available or zero otherwise.
  • panels: 1-D array of ids for panel formation, where each choice situation has a panel for each variable
  • set_vars: Dictionary specifying some variables and the values that they are to be set
  • include_correlations: Boolean whether or not to consider correlation between variables
  • optim_method: String representing with optimisation method to use. Options are "L-BFGS-scipy", "BFGS-scipy", "L-BFGS-jax", and "BFGS-jax". Scipy uses the standard scipy library and jax uses jax's scipy library, which is significantly faster (may not work with batching) but may not be maintained and potentially discontinued without notice.
importpandasaspdimportnumpyasnpimportjaxfromjaxlogit.mixed_logitimportMixedLogit, ConfigData# 64bit precisionjax.config.update("jax_enable_x64", True)
# read and format datadf_wide=pd.read_table("http://transp-or.epfl.ch/data/swissmetro.dat", sep='\t')
# Keep only observations for commute and business purposes that contain known choicesdf_wide=df_wide[(df_wide['PURPOSE'].isin([1, 3]) & (df_wide['CHOICE'] !=0))]
df_wide['custom_id'] =np.arange(len(df_wide)) # Add unique identifierdf_wide['CHOICE'] =df_wide['CHOICE'].map({1: 'TRAIN', 2:'SM', 3: 'CAR'})
fromjaxlogit.utilsimportwide_to_longdf=wide_to_long(df_wide, id_col='custom_id', alt_name='alt', sep='_',
alt_list=['TRAIN', 'SM', 'CAR'], empty_val=0,
varying=['TT', 'CO', 'HE', 'AV', 'SEATS'], alt_is_prefix=True)
# modification of data to includedf['ASC_TRAIN'] =np.ones(len(df))*(df['alt'] =='TRAIN')
df['ASC_CAR'] =np.ones(len(df))*(df['alt'] =='CAR')
df['TT'], df['CO'] =df['TT']/100, df['CO']/100# Scale variablesannual_pass= (df['GA'] ==1) & (df['alt'].isin(['TRAIN', 'SM']))
df.loc[annual_pass, 'CO'] =0# Cost zero for pass holders# specification of variablesvarnames=['ASC_CAR', 'ASC_TRAIN', 'ASC_SM', 'CO', 'TT']
randvars={'ASC_CAR': 'n', 'ASC_TRAIN': 'n', 'ASC_SM': 'n'}
set_vars= {
'ASC_SM': 0.0, 'sd.ASC_TRAIN': 1.0, 'sd.ASC_CAR': 0.0,
'chol.ASC_CAR.ASC_TRAIN': 0.0, 'chol.ASC_CAR.ASC_SM': 0.0
} # Identification of error components, see J. Walker's PhD thesis (MIT 2001)# change some optional argumentsconfig=ConfigData(
avail=df['AV'],
panels=df["ID"],
set_vars=set_vars,
include_correlations=True, # Enable correlation between random parametersoptim_method="L-BFGS-scipy"
)
model=MixedLogit()
res=model.fit(
X=df[varnames],
y=df['CHOICE'],
varnames=varnames,
alts=df['alt'],
ids=df['custom_id'],
randvars=randvars,
config=config
)
model.summary()
 Message: CONVERGENCE: RELATIVE REDUCTION OF F <= FACTR*EPSMCH
Iterations: 17
Function evaluations: 21
Estimation time= 57.3 seconds
---------------------------------------------------------------------------
Coefficient Estimate Std.Err. z-val P>|z|
---------------------------------------------------------------------------
ASC_CAR -0.2279764 0.4655508 -0.4896917 0.624 ASC_TRAIN -1.1966263 0.5873328 -2.0373905 0.0416 * ASC_SM 0.1000000 0.0000000 inf 0 ***
CO -2.0490264 0.3444999 -5.9478289 2.85e-09 ***
TT -2.1429727 0.6607183 -3.2433986 0.00119 ** sd.ASC_CAR 0.1000000 0.0000000 inf 0 ***
sd.ASC_TRAIN 0.1000000 0.0000000 inf 0 ***
sd.ASC_SM 2.4459162 0.2814464 8.6905216 4.47e-18 ***
chol.ASC_CAR.ASC_TR 0.1000000 0.0000000 inf 0 ***
chol.ASC_CAR.ASC_SM 0.1000000 0.0000000 inf 0 ***
chol.ASC_TRAIN.ASC_ 0.6937107 0.2260149 3.0693135 0.00215 ** ---------------------------------------------------------------------------
Significance: 0 '***' 0.001 '**' 0.01 '*' 0.05 '.' 0.1 ' ' 1
Log-Likelihood= -4071.438
AIC= 8154.876
BIC= 8195.795

If out of memory, the data can be batched as well.

Quick Install

Install jaxlogit using pip: pip install jaxlogit.

Alternatively, clone the repo.

Benchmark

As shown in the plot below, jaxlogit with batching uses significantly less memory than xlogit. Graph comparing memory usage of jaxlogit and xlogit

Graph comparing time of jaxlogit, xlogit, and biogeme

The graph of memory shows that using the batching reduces memory usage below that of xlogit and other types of jaxlogit. Off the edge of this graph is a spike in biogeme's memory usage, up to 30GB.

In the timing graph, it can be seen that batching does not have a significant impact on time taken. Additionally, the jaxlogit performs much faster when using the experimental jax methods. The time for standard scipy method and xlogit are quite comparable. Biogeme is slower than xlogit and any form of jaxlogit.

These were taken on the complicated electricity dataset. On the nicer artificial dataset the timing looks like this: Graph comparing timing of jaxlogit and xlogit

About

No description, website, or topics provided.

Resources

Stars

2 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages

, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Auto-enable theater mode on YouTube\n(function() {\n function tryTheater() {\n var btn = document.querySelector('button[aria-label=\"Theater mode\"], ytd-player #player button[title=\"Theater mode\"]');\n if (btn && !btn.classList.contains('activated')) {\n btn.click();\n }\n }\n \n // Try immediately\n tryTheater();\n \n // Try after navigation (SPA)\n var lastUrl = location.href;\n setInterval(function() {\n if (location.href !== lastUrl) {\n lastUrl = location.href;\n setTimeout(tryTheater, 500);\n }\n }, 1000);\n \n // Also try on player load\n var observer = new MutationObserver(tryTheater);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "YouTube Theater Mode Default"); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content

Repository files navigation

Jaxlogit

MSLE estimation of linear-in-parameters mixed logit models. This package is based on xlogit, with the core computational engine replaced by jax. Additionally, the documentation structure and many examples are adapted from xlogit's documentation.

Why use Jaxlogit?

While there are many packages out there that can compute mixed logit models, jaxlogit is unique in its processing speed and capabilities, which come from using jax and just in time compilation. You can see the comparison here. Because jaxlogit specialises in linear models, it can use accelerated linear algebra to do the matrix multiplication, which significantly increases the speed. Jaxlogit can also work with large quanities of data by calculating them in batches, reducing memory load. This can be seen in this example.

Quick Start

This example uses jaxlogit to to estimate a mixed logit model for choices of transport method using the swissmetro dataset. It contains stated-preferences for three alternative transportation modes that include car, train and a newly introduced mode: the swissmetro. The dataset is available here and Bierlaire et. al., (2001) provides a detailed discussion of the data as well as its context and collection process.

The explanatory variables are cost and travel time of each of the alternatives, as well as alternative-specific constraints to represent unobserved factors. New functionality used compared to xlogit is considering the correlation between the distributions of explanatory variables. In this example, only correlations between some variables are considered, and some are set to certain values, according to research done in J. Walker's PhD thesis (MIT 2001). This is specified by specifying some correlation parameters to be 0.

The data must be in long format. Wide format data can be transformed into long format data using the wide_to_long function.

Required parameters:

  • X: 2-D array of input data (in long format) with choice situations as rows, and variables as columns
  • y: 1-D array of choices (in long format)
  • varnames: List of variable names that matches the number and order of the columns in X
  • alts: 1-D array of alternative representing the alternative chosen for each input
  • ids: 1-D array of the ids of the choice situations
  • randvars: dictionary of variables and their mixing distributions. In this example, only "n" normal variables are used. Valid distributions are "n" normal, "ln" log normal and "n_trunc" truncated normal.

Optional Parameters in this example:

  • avail: 2-D array of availability of alternatives for the choice situations. One when available or zero otherwise.
  • panels: 1-D array of ids for panel formation, where each choice situation has a panel for each variable
  • set_vars: Dictionary specifying some variables and the values that they are to be set
  • include_correlations: Boolean whether or not to consider correlation between variables
  • optim_method: String representing with optimisation method to use. Options are "L-BFGS-scipy", "BFGS-scipy", "L-BFGS-jax", and "BFGS-jax". Scipy uses the standard scipy library and jax uses jax's scipy library, which is significantly faster (may not work with batching) but may not be maintained and potentially discontinued without notice.
importpandasaspdimportnumpyasnpimportjaxfromjaxlogit.mixed_logitimportMixedLogit, ConfigData# 64bit precisionjax.config.update("jax_enable_x64", True)
# read and format datadf_wide=pd.read_table("http://transp-or.epfl.ch/data/swissmetro.dat", sep='\t')
# Keep only observations for commute and business purposes that contain known choicesdf_wide=df_wide[(df_wide['PURPOSE'].isin([1, 3]) & (df_wide['CHOICE'] !=0))]
df_wide['custom_id'] =np.arange(len(df_wide)) # Add unique identifierdf_wide['CHOICE'] =df_wide['CHOICE'].map({1: 'TRAIN', 2:'SM', 3: 'CAR'})
fromjaxlogit.utilsimportwide_to_longdf=wide_to_long(df_wide, id_col='custom_id', alt_name='alt', sep='_',
alt_list=['TRAIN', 'SM', 'CAR'], empty_val=0,
varying=['TT', 'CO', 'HE', 'AV', 'SEATS'], alt_is_prefix=True)
# modification of data to includedf['ASC_TRAIN'] =np.ones(len(df))*(df['alt'] =='TRAIN')
df['ASC_CAR'] =np.ones(len(df))*(df['alt'] =='CAR')
df['TT'], df['CO'] =df['TT']/100, df['CO']/100# Scale variablesannual_pass= (df['GA'] ==1) & (df['alt'].isin(['TRAIN', 'SM']))
df.loc[annual_pass, 'CO'] =0# Cost zero for pass holders# specification of variablesvarnames=['ASC_CAR', 'ASC_TRAIN', 'ASC_SM', 'CO', 'TT']
randvars={'ASC_CAR': 'n', 'ASC_TRAIN': 'n', 'ASC_SM': 'n'}
set_vars= {
'ASC_SM': 0.0, 'sd.ASC_TRAIN': 1.0, 'sd.ASC_CAR': 0.0,
'chol.ASC_CAR.ASC_TRAIN': 0.0, 'chol.ASC_CAR.ASC_SM': 0.0
} # Identification of error components, see J. Walker's PhD thesis (MIT 2001)# change some optional argumentsconfig=ConfigData(
avail=df['AV'],
panels=df["ID"],
set_vars=set_vars,
include_correlations=True, # Enable correlation between random parametersoptim_method="L-BFGS-scipy"
)
model=MixedLogit()
res=model.fit(
X=df[varnames],
y=df['CHOICE'],
varnames=varnames,
alts=df['alt'],
ids=df['custom_id'],
randvars=randvars,
config=config
)
model.summary()
 Message: CONVERGENCE: RELATIVE REDUCTION OF F <= FACTR*EPSMCH
Iterations: 17
Function evaluations: 21
Estimation time= 57.3 seconds
---------------------------------------------------------------------------
Coefficient Estimate Std.Err. z-val P>|z|
---------------------------------------------------------------------------
ASC_CAR -0.2279764 0.4655508 -0.4896917 0.624 ASC_TRAIN -1.1966263 0.5873328 -2.0373905 0.0416 * ASC_SM 0.1000000 0.0000000 inf 0 ***
CO -2.0490264 0.3444999 -5.9478289 2.85e-09 ***
TT -2.1429727 0.6607183 -3.2433986 0.00119 ** sd.ASC_CAR 0.1000000 0.0000000 inf 0 ***
sd.ASC_TRAIN 0.1000000 0.0000000 inf 0 ***
sd.ASC_SM 2.4459162 0.2814464 8.6905216 4.47e-18 ***
chol.ASC_CAR.ASC_TR 0.1000000 0.0000000 inf 0 ***
chol.ASC_CAR.ASC_SM 0.1000000 0.0000000 inf 0 ***
chol.ASC_TRAIN.ASC_ 0.6937107 0.2260149 3.0693135 0.00215 ** ---------------------------------------------------------------------------
Significance: 0 '***' 0.001 '**' 0.01 '*' 0.05 '.' 0.1 ' ' 1
Log-Likelihood= -4071.438
AIC= 8154.876
BIC= 8195.795

If out of memory, the data can be batched as well.

Quick Install

Install jaxlogit using pip: pip install jaxlogit.

Alternatively, clone the repo.

Benchmark

As shown in the plot below, jaxlogit with batching uses significantly less memory than xlogit. Graph comparing memory usage of jaxlogit and xlogit

Graph comparing time of jaxlogit, xlogit, and biogeme

The graph of memory shows that using the batching reduces memory usage below that of xlogit and other types of jaxlogit. Off the edge of this graph is a spike in biogeme's memory usage, up to 30GB.

In the timing graph, it can be seen that batching does not have a significant impact on time taken. Additionally, the jaxlogit performs much faster when using the experimental jax methods. The time for standard scipy method and xlogit are quite comparable. Biogeme is slower than xlogit and any form of jaxlogit.

These were taken on the complicated electricity dataset. On the nicer artificial dataset the timing looks like this: Graph comparing timing of jaxlogit and xlogit

About

No description, website, or topics provided.

Resources

Stars

2 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages

, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Remove or un-stick sticky/fixed headers that block content\n(function() {\n function unstick() {\n document.querySelectorAll('header, nav, [role=\"banner\"], .header, .navbar, .sticky, .fixed-top, [style*=\"position: fixed\"], [style*=\"position:sticky\"]').forEach(function(el) {\n if (el.style.position === 'fixed' || el.style.position === 'sticky' || \n getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') {\n el.style.position = 'static';\n el.style.top = 'auto';\n el.style.zIndex = 'auto';\n }\n });\n }\n \n unstick();\n \n var observer = new MutationObserver(unstick);\n observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] });\n})();", "Kill Sticky Headers"); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content

Repository files navigation

Jaxlogit

MSLE estimation of linear-in-parameters mixed logit models. This package is based on xlogit, with the core computational engine replaced by jax. Additionally, the documentation structure and many examples are adapted from xlogit's documentation.

Why use Jaxlogit?

While there are many packages out there that can compute mixed logit models, jaxlogit is unique in its processing speed and capabilities, which come from using jax and just in time compilation. You can see the comparison here. Because jaxlogit specialises in linear models, it can use accelerated linear algebra to do the matrix multiplication, which significantly increases the speed. Jaxlogit can also work with large quanities of data by calculating them in batches, reducing memory load. This can be seen in this example.

Quick Start

This example uses jaxlogit to to estimate a mixed logit model for choices of transport method using the swissmetro dataset. It contains stated-preferences for three alternative transportation modes that include car, train and a newly introduced mode: the swissmetro. The dataset is available here and Bierlaire et. al., (2001) provides a detailed discussion of the data as well as its context and collection process.

The explanatory variables are cost and travel time of each of the alternatives, as well as alternative-specific constraints to represent unobserved factors. New functionality used compared to xlogit is considering the correlation between the distributions of explanatory variables. In this example, only correlations between some variables are considered, and some are set to certain values, according to research done in J. Walker's PhD thesis (MIT 2001). This is specified by specifying some correlation parameters to be 0.

The data must be in long format. Wide format data can be transformed into long format data using the wide_to_long function.

Required parameters:

  • X: 2-D array of input data (in long format) with choice situations as rows, and variables as columns
  • y: 1-D array of choices (in long format)
  • varnames: List of variable names that matches the number and order of the columns in X
  • alts: 1-D array of alternative representing the alternative chosen for each input
  • ids: 1-D array of the ids of the choice situations
  • randvars: dictionary of variables and their mixing distributions. In this example, only "n" normal variables are used. Valid distributions are "n" normal, "ln" log normal and "n_trunc" truncated normal.

Optional Parameters in this example:

  • avail: 2-D array of availability of alternatives for the choice situations. One when available or zero otherwise.
  • panels: 1-D array of ids for panel formation, where each choice situation has a panel for each variable
  • set_vars: Dictionary specifying some variables and the values that they are to be set
  • include_correlations: Boolean whether or not to consider correlation between variables
  • optim_method: String representing with optimisation method to use. Options are "L-BFGS-scipy", "BFGS-scipy", "L-BFGS-jax", and "BFGS-jax". Scipy uses the standard scipy library and jax uses jax's scipy library, which is significantly faster (may not work with batching) but may not be maintained and potentially discontinued without notice.
importpandasaspdimportnumpyasnpimportjaxfromjaxlogit.mixed_logitimportMixedLogit, ConfigData# 64bit precisionjax.config.update("jax_enable_x64", True)
# read and format datadf_wide=pd.read_table("http://transp-or.epfl.ch/data/swissmetro.dat", sep='\t')
# Keep only observations for commute and business purposes that contain known choicesdf_wide=df_wide[(df_wide['PURPOSE'].isin([1, 3]) & (df_wide['CHOICE'] !=0))]
df_wide['custom_id'] =np.arange(len(df_wide)) # Add unique identifierdf_wide['CHOICE'] =df_wide['CHOICE'].map({1: 'TRAIN', 2:'SM', 3: 'CAR'})
fromjaxlogit.utilsimportwide_to_longdf=wide_to_long(df_wide, id_col='custom_id', alt_name='alt', sep='_',
alt_list=['TRAIN', 'SM', 'CAR'], empty_val=0,
varying=['TT', 'CO', 'HE', 'AV', 'SEATS'], alt_is_prefix=True)
# modification of data to includedf['ASC_TRAIN'] =np.ones(len(df))*(df['alt'] =='TRAIN')
df['ASC_CAR'] =np.ones(len(df))*(df['alt'] =='CAR')
df['TT'], df['CO'] =df['TT']/100, df['CO']/100# Scale variablesannual_pass= (df['GA'] ==1) & (df['alt'].isin(['TRAIN', 'SM']))
df.loc[annual_pass, 'CO'] =0# Cost zero for pass holders# specification of variablesvarnames=['ASC_CAR', 'ASC_TRAIN', 'ASC_SM', 'CO', 'TT']
randvars={'ASC_CAR': 'n', 'ASC_TRAIN': 'n', 'ASC_SM': 'n'}
set_vars= {
'ASC_SM': 0.0, 'sd.ASC_TRAIN': 1.0, 'sd.ASC_CAR': 0.0,
'chol.ASC_CAR.ASC_TRAIN': 0.0, 'chol.ASC_CAR.ASC_SM': 0.0
} # Identification of error components, see J. Walker's PhD thesis (MIT 2001)# change some optional argumentsconfig=ConfigData(
avail=df['AV'],
panels=df["ID"],
set_vars=set_vars,
include_correlations=True, # Enable correlation between random parametersoptim_method="L-BFGS-scipy"
)
model=MixedLogit()
res=model.fit(
X=df[varnames],
y=df['CHOICE'],
varnames=varnames,
alts=df['alt'],
ids=df['custom_id'],
randvars=randvars,
config=config
)
model.summary()
 Message: CONVERGENCE: RELATIVE REDUCTION OF F <= FACTR*EPSMCH
Iterations: 17
Function evaluations: 21
Estimation time= 57.3 seconds
---------------------------------------------------------------------------
Coefficient Estimate Std.Err. z-val P>|z|
---------------------------------------------------------------------------
ASC_CAR -0.2279764 0.4655508 -0.4896917 0.624 ASC_TRAIN -1.1966263 0.5873328 -2.0373905 0.0416 * ASC_SM 0.1000000 0.0000000 inf 0 ***
CO -2.0490264 0.3444999 -5.9478289 2.85e-09 ***
TT -2.1429727 0.6607183 -3.2433986 0.00119 ** sd.ASC_CAR 0.1000000 0.0000000 inf 0 ***
sd.ASC_TRAIN 0.1000000 0.0000000 inf 0 ***
sd.ASC_SM 2.4459162 0.2814464 8.6905216 4.47e-18 ***
chol.ASC_CAR.ASC_TR 0.1000000 0.0000000 inf 0 ***
chol.ASC_CAR.ASC_SM 0.1000000 0.0000000 inf 0 ***
chol.ASC_TRAIN.ASC_ 0.6937107 0.2260149 3.0693135 0.00215 ** ---------------------------------------------------------------------------
Significance: 0 '***' 0.001 '**' 0.01 '*' 0.05 '.' 0.1 ' ' 1
Log-Likelihood= -4071.438
AIC= 8154.876
BIC= 8195.795

If out of memory, the data can be batched as well.

Quick Install

Install jaxlogit using pip: pip install jaxlogit.

Alternatively, clone the repo.

Benchmark

As shown in the plot below, jaxlogit with batching uses significantly less memory than xlogit. Graph comparing memory usage of jaxlogit and xlogit

Graph comparing time of jaxlogit, xlogit, and biogeme

The graph of memory shows that using the batching reduces memory usage below that of xlogit and other types of jaxlogit. Off the edge of this graph is a spike in biogeme's memory usage, up to 30GB.

In the timing graph, it can be seen that batching does not have a significant impact on time taken. Additionally, the jaxlogit performs much faster when using the experimental jax methods. The time for standard scipy method and xlogit are quite comparable. Biogeme is slower than xlogit and any form of jaxlogit.

These were taken on the complicated electricity dataset. On the nicer artificial dataset the timing looks like this: Graph comparing timing of jaxlogit and xlogit

About

No description, website, or topics provided.

Resources

Stars

2 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages

, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Universal Dark Mode - works on any site\n(function() {\n var enabled = true;\n \n function applyDarkMode() {\n if (!enabled) return;\n \n // Create style element if it doesn't exist\n var style = document.getElementById('universal-dark-mode-style');\n if (!style) {\n style = document.createElement('style');\n style.id = 'universal-dark-mode-style';\n document.head.appendChild(style);\n }\n \n // Dark mode CSS - inverts colors but preserves images/video\n style.textContent = '\n /* Invert everything except media */\n html {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #1a1a2e !important;\n }\n \n /* Restore images, videos, iframes, canvas */\n img, video, iframe, canvas, svg, picture, [style*=\"background-image\"] {\n filter: invert(1) hue-rotate(180deg) !important;\n }\n \n /* Preserve specific elements that should not be inverted */\n .no-dark-mode, .no-dark-mode *,\n [data-theme=\"light\"], [data-theme=\"light\"],\n .ace_editor, .ace_editor *,\n .CodeMirror, .CodeMirror *,\n .monaco-editor, .monaco-editor *,\n .markdown-body pre, .markdown-body pre *,\n .highlight, .highlight *,\n pre code, pre code * {\n filter: none !important;\n }\n \n /* Fix common UI elements */\n .modal, .popup, .dropdown-menu, .tooltip, .popover {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #2d2d44 !important;\n border-color: #444 !important;\n }\n \n /* Scrollbars */\n ::-webkit-scrollbar { background: #1a1a2e !important; }\n ::-webkit-scrollbar-thumb { background: #444 !important; }\n ::-webkit-scrollbar-thumb:hover { background: #555 !important; }\n \n /* Selection */\n ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ';\n }\n \n function removeDarkMode() {\n var style = document.getElementById('universal-dark-mode-style');\n if (style) style.remove();\n }\n \n // Toggle with Alt+Shift+D\n document.addEventListener('keydown', function(e) {\n if (e.altKey && e.shiftKey && e.key === 'D') {\n e.preventDefault();\n enabled = !enabled;\n if (enabled) {\n applyDarkMode();\n console.log('[Universal Dark Mode] Enabled');\n } else {\n removeDarkMode();\n console.log('[Universal Dark Mode] Disabled');\n }\n }\n });\n \n // Apply on load\n applyDarkMode();\n \n // Re-apply on dynamic content\n var observer = new MutationObserver(function(mutations) {\n if (enabled && !document.getElementById('universal-dark-mode-style')) {\n applyDarkMode();\n }\n });\n observer.observe(document.head, { childList: true });\n \n console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle');\n})();", "Universal Dark Mode"); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })();
Skip to content

Repository files navigation

Jaxlogit

MSLE estimation of linear-in-parameters mixed logit models. This package is based on xlogit, with the core computational engine replaced by jax. Additionally, the documentation structure and many examples are adapted from xlogit's documentation.

Why use Jaxlogit?

While there are many packages out there that can compute mixed logit models, jaxlogit is unique in its processing speed and capabilities, which come from using jax and just in time compilation. You can see the comparison here. Because jaxlogit specialises in linear models, it can use accelerated linear algebra to do the matrix multiplication, which significantly increases the speed. Jaxlogit can also work with large quanities of data by calculating them in batches, reducing memory load. This can be seen in this example.

Quick Start

This example uses jaxlogit to to estimate a mixed logit model for choices of transport method using the swissmetro dataset. It contains stated-preferences for three alternative transportation modes that include car, train and a newly introduced mode: the swissmetro. The dataset is available here and Bierlaire et. al., (2001) provides a detailed discussion of the data as well as its context and collection process.

The explanatory variables are cost and travel time of each of the alternatives, as well as alternative-specific constraints to represent unobserved factors. New functionality used compared to xlogit is considering the correlation between the distributions of explanatory variables. In this example, only correlations between some variables are considered, and some are set to certain values, according to research done in J. Walker's PhD thesis (MIT 2001). This is specified by specifying some correlation parameters to be 0.

The data must be in long format. Wide format data can be transformed into long format data using the wide_to_long function.

Required parameters:

  • X: 2-D array of input data (in long format) with choice situations as rows, and variables as columns
  • y: 1-D array of choices (in long format)
  • varnames: List of variable names that matches the number and order of the columns in X
  • alts: 1-D array of alternative representing the alternative chosen for each input
  • ids: 1-D array of the ids of the choice situations
  • randvars: dictionary of variables and their mixing distributions. In this example, only "n" normal variables are used. Valid distributions are "n" normal, "ln" log normal and "n_trunc" truncated normal.

Optional Parameters in this example:

  • avail: 2-D array of availability of alternatives for the choice situations. One when available or zero otherwise.
  • panels: 1-D array of ids for panel formation, where each choice situation has a panel for each variable
  • set_vars: Dictionary specifying some variables and the values that they are to be set
  • include_correlations: Boolean whether or not to consider correlation between variables
  • optim_method: String representing with optimisation method to use. Options are "L-BFGS-scipy", "BFGS-scipy", "L-BFGS-jax", and "BFGS-jax". Scipy uses the standard scipy library and jax uses jax's scipy library, which is significantly faster (may not work with batching) but may not be maintained and potentially discontinued without notice.
importpandasaspdimportnumpyasnpimportjaxfromjaxlogit.mixed_logitimportMixedLogit, ConfigData# 64bit precisionjax.config.update("jax_enable_x64", True)
# read and format datadf_wide=pd.read_table("http://transp-or.epfl.ch/data/swissmetro.dat", sep='\t')
# Keep only observations for commute and business purposes that contain known choicesdf_wide=df_wide[(df_wide['PURPOSE'].isin([1, 3]) & (df_wide['CHOICE'] !=0))]
df_wide['custom_id'] =np.arange(len(df_wide)) # Add unique identifierdf_wide['CHOICE'] =df_wide['CHOICE'].map({1: 'TRAIN', 2:'SM', 3: 'CAR'})
fromjaxlogit.utilsimportwide_to_longdf=wide_to_long(df_wide, id_col='custom_id', alt_name='alt', sep='_',
alt_list=['TRAIN', 'SM', 'CAR'], empty_val=0,
varying=['TT', 'CO', 'HE', 'AV', 'SEATS'], alt_is_prefix=True)
# modification of data to includedf['ASC_TRAIN'] =np.ones(len(df))*(df['alt'] =='TRAIN')
df['ASC_CAR'] =np.ones(len(df))*(df['alt'] =='CAR')
df['TT'], df['CO'] =df['TT']/100, df['CO']/100# Scale variablesannual_pass= (df['GA'] ==1) & (df['alt'].isin(['TRAIN', 'SM']))
df.loc[annual_pass, 'CO'] =0# Cost zero for pass holders# specification of variablesvarnames=['ASC_CAR', 'ASC_TRAIN', 'ASC_SM', 'CO', 'TT']
randvars={'ASC_CAR': 'n', 'ASC_TRAIN': 'n', 'ASC_SM': 'n'}
set_vars= {
'ASC_SM': 0.0, 'sd.ASC_TRAIN': 1.0, 'sd.ASC_CAR': 0.0,
'chol.ASC_CAR.ASC_TRAIN': 0.0, 'chol.ASC_CAR.ASC_SM': 0.0
} # Identification of error components, see J. Walker's PhD thesis (MIT 2001)# change some optional argumentsconfig=ConfigData(
avail=df['AV'],
panels=df["ID"],
set_vars=set_vars,
include_correlations=True, # Enable correlation between random parametersoptim_method="L-BFGS-scipy"
)
model=MixedLogit()
res=model.fit(
X=df[varnames],
y=df['CHOICE'],
varnames=varnames,
alts=df['alt'],
ids=df['custom_id'],
randvars=randvars,
config=config
)
model.summary()
 Message: CONVERGENCE: RELATIVE REDUCTION OF F <= FACTR*EPSMCH
Iterations: 17
Function evaluations: 21
Estimation time= 57.3 seconds
---------------------------------------------------------------------------
Coefficient Estimate Std.Err. z-val P>|z|
---------------------------------------------------------------------------
ASC_CAR -0.2279764 0.4655508 -0.4896917 0.624 ASC_TRAIN -1.1966263 0.5873328 -2.0373905 0.0416 * ASC_SM 0.1000000 0.0000000 inf 0 ***
CO -2.0490264 0.3444999 -5.9478289 2.85e-09 ***
TT -2.1429727 0.6607183 -3.2433986 0.00119 ** sd.ASC_CAR 0.1000000 0.0000000 inf 0 ***
sd.ASC_TRAIN 0.1000000 0.0000000 inf 0 ***
sd.ASC_SM 2.4459162 0.2814464 8.6905216 4.47e-18 ***
chol.ASC_CAR.ASC_TR 0.1000000 0.0000000 inf 0 ***
chol.ASC_CAR.ASC_SM 0.1000000 0.0000000 inf 0 ***
chol.ASC_TRAIN.ASC_ 0.6937107 0.2260149 3.0693135 0.00215 ** ---------------------------------------------------------------------------
Significance: 0 '***' 0.001 '**' 0.01 '*' 0.05 '.' 0.1 ' ' 1
Log-Likelihood= -4071.438
AIC= 8154.876
BIC= 8195.795

If out of memory, the data can be batched as well.

Quick Install

Install jaxlogit using pip: pip install jaxlogit.

Alternatively, clone the repo.

Benchmark

As shown in the plot below, jaxlogit with batching uses significantly less memory than xlogit. Graph comparing memory usage of jaxlogit and xlogit

Graph comparing time of jaxlogit, xlogit, and biogeme

The graph of memory shows that using the batching reduces memory usage below that of xlogit and other types of jaxlogit. Off the edge of this graph is a spike in biogeme's memory usage, up to 30GB.

In the timing graph, it can be seen that batching does not have a significant impact on time taken. Additionally, the jaxlogit performs much faster when using the experimental jax methods. The time for standard scipy method and xlogit are quite comparable. Biogeme is slower than xlogit and any form of jaxlogit.

These were taken on the complicated electricity dataset. On the nicer artificial dataset the timing looks like this: Graph comparing timing of jaxlogit and xlogit

About

No description, website, or topics provided.

Resources

Stars

2 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages