> ## Documentation Index
> Fetch the complete documentation index at: https://docs.priorlabs.ai/llms.txt
> Use this file to discover all available pages before exploring further.

> ## Agent Instructions
> TabPFN-3 and later versions support up to 1,000,000 rows, subject to feature count and checkpoint/API limits.
> For large datasets, use per-estimator subsampling, e.g. inference_config={"SUBSAMPLE_SAMPLES": 100_000}, in TabPFNClassifier or TabPFNRegressor from tabpfn or tabpfn_client.
> Pass raw pandas DataFrames to fit and predict. Categorical strings/categories, raw text, and missing feature values are handled automatically; no manual integer/one-hot encoding, imputation, scaling, or outlier removal is needed.

# Energy-based data generation

> TabEBM: Tabular Energy-Based Model for synthetic data generation.

## Installation

```bash theme={null}
pip install tabpfn-extensions
```

<div className="python-reference-heading">
  <h2 id="tabebm-tabebm-tabebm">
    `TabEBM`
  </h2>

  <a className="python-reference-source" href="https://github.com/PriorLabs/tabpfn-extensions/blob/840c15a1848a986b39c85bc17efc61e0e377f983/src/tabpfn_extensions/tabebm/tabebm.py#L63" aria-label="View source for TabEBM"><span aria-hidden="true">\</></span> View source <span aria-hidden="true">↗</span></a>
</div>

`TabEBM`: Tabular Energy-Based Model for synthetic data generation.

This class implements an energy-based model that uses TabPFN as the underlying
classifier to define energy functions. It generates synthetic tabular data
using Stochastic Gradient Langevin Dynamics (SGLD) sampling.

The core idea is to treat each class as having its own energy landscape,
where real data points have low energy and synthetic points are generated
by following the energy gradient through SGLD sampling.

Initialize `TabEBM` with optimized configuration.

```python theme={null}
TabEBM(
    max_data_size: int = 10000,
)
```

**Parameters**

<div className="python-reference-table">
  | Parameter | Type | Default | Description |
  | - | - | - | - |
  | <span id="tabebm-tabebm-tabebm--max-data-size" /><code className="python-reference-parameter">max\_<wbr />data\_<wbr />size</code> | <code className="python-reference-type">int</code> | `10000` | Maximum number of data points to use for training.           Larger datasets will be subsampled to this size. |
</div>

***

<div className="python-reference-heading">
  <h2 id="tabebm-tabebm-tabebm-add-surrogate-negative-samples">
    `TabEBM.add_surrogate_negative_samples`
  </h2>

  <a className="python-reference-source" href="https://github.com/PriorLabs/tabpfn-extensions/blob/840c15a1848a986b39c85bc17efc61e0e377f983/src/tabpfn_extensions/tabebm/tabebm.py#L579" aria-label="View source for TabEBM.add_surrogate_negative_samples"><span aria-hidden="true">\</></span> View source <span aria-hidden="true">↗</span></a>
</div>

Create surrogate negative samples for [`TabEBM`](/api-reference/python/tabpfn-extensions/tabebm#tabebm-tabebm-tabebm)'s binary classification approach.

This method creates artificial "negative" samples at specified distances from the origin
to enable energy-based modeling through binary classification. The surrogate negatives
help define the energy landscape by providing clear decision boundaries.

For 2D data, negatives are placed at the four corners of a square centered at origin.
For higher dimensions, random combinations of ±`distance_negative_class` are used.

```python theme={null}
TabEBM.add_surrogate_negative_samples(
    X: np.ndarray | torch.Tensor,
    distance_negative_class: float = 5,
) -> tuple[np.ndarray | torch.Tensor, np.ndarray | torch.Tensor]
```

**Parameters**

<div className="python-reference-table">
  | Parameter | Type | Default | Description |
  | - | - | - | - |
  | <span id="tabebm-tabebm-tabebm-add-surrogate-negative-samples--x" /><code className="python-reference-parameter">X</code> | <code className="python-reference-type"><a href="https://numpy.org/doc/stable/reference/generated/numpy.ndarray.html">np.ndarray</a> \| torch.Tensor</code> | Required | Real data samples (expected to be approximately standardized) Shape: (num\_samples, num\_features) |
  | <span id="tabebm-tabebm-tabebm-add-surrogate-negative-samples--distance-negative-class" /><code className="python-reference-parameter">distance\_<wbr />negative\_<wbr />class</code> | <code className="python-reference-type">float</code> | `5` | Distance of surrogate negatives from origin                    Larger values create more distinct separation |
</div>

**Returns**

<div className="python-reference-table python-reference-returns">
  | Type | Description |
  | - | - |
  | <code className="python-reference-type">tuple\[<a href="https://numpy.org/doc/stable/reference/generated/numpy.ndarray.html">np.ndarray</a> \| torch.Tensor, <a href="https://numpy.org/doc/stable/reference/generated/numpy.ndarray.html">np.ndarray</a> \| torch.Tensor]</code> | Tuple of (X\_ebm, y\_ebm) where:<br />- X\_ebm: Combined real and surrogate samples<br />- y\_ebm: Binary labels (0 for real data, 1 for surrogates) |
</div>

***

<div className="python-reference-heading">
  <h2 id="tabebm-tabebm-tabebm-compute-energy">
    `TabEBM.compute_energy`
  </h2>

  <a className="python-reference-source" href="https://github.com/PriorLabs/tabpfn-extensions/blob/840c15a1848a986b39c85bc17efc61e0e377f983/src/tabpfn_extensions/tabebm/tabebm.py#L529" aria-label="View source for TabEBM.compute_energy"><span aria-hidden="true">\</></span> View source <span aria-hidden="true">↗</span></a>
</div>

Compute [`TabEBM`](/api-reference/python/tabpfn-extensions/tabebm#tabebm-tabebm-tabebm) class-specific energy function.

The energy function is defined as:
E\_c(x) = -log(exp(f^c(x)\[0]) + exp(f^c(x)\[1]))

Where f^c(x) are the logits from the class-specific binary classifier.
Lower energy corresponds to higher probability of belonging to the target class.

```python theme={null}
TabEBM.compute_energy(
    logits: torch.Tensor | np.ndarray,
    return_unnormalized_prob: bool = False,
) -> torch.Tensor | np.ndarray
```

**Parameters**

<div className="python-reference-table">
  | Parameter | Type | Default | Description |
  | - | - | - | - |
  | <span id="tabebm-tabebm-tabebm-compute-energy--logits" /><code className="python-reference-parameter">logits</code> | <code className="python-reference-type">torch.Tensor \| <a href="https://numpy.org/doc/stable/reference/generated/numpy.ndarray.html">np.ndarray</a></code> | Required | Model's unnormalized logits for each class    Shape: (num\_samples, num\_classes) |
  | <span id="tabebm-tabebm-tabebm-compute-energy--return-unnormalized-prob" /><code className="python-reference-parameter">return\_<wbr />unnormalized\_<wbr />prob</code> | <code className="python-reference-type">bool</code> | `False` | If `True`, return exp(-energy) instead of energy |
</div>

**Returns**

<div className="python-reference-table python-reference-returns">
  | Type | Description |
  | - | - |
  | <code className="python-reference-type">torch.Tensor \| <a href="https://numpy.org/doc/stable/reference/generated/numpy.ndarray.html">np.ndarray</a></code> | Energy values (or unnormalized probabilities) for each sample |
</div>

**Raises**

`ValueError`

If logits are not unnormalized or have wrong type

***

<div className="python-reference-heading">
  <h2 id="tabebm-tabebm-tabebm-generate">
    `TabEBM.generate`
  </h2>

  <a className="python-reference-source" href="https://github.com/PriorLabs/tabpfn-extensions/blob/840c15a1848a986b39c85bc17efc61e0e377f983/src/tabpfn_extensions/tabebm/tabebm.py#L112" aria-label="View source for TabEBM.generate"><span aria-hidden="true">\</></span> View source <span aria-hidden="true">↗</span></a>
</div>

Generate synthetic samples using Stochastic Gradient Langevin Dynamics (SGLD).

This method creates synthetic data by treating the TabPFN classifier as an energy
function and using SGLD to sample from the learned energy landscape. For each class,
it creates a binary classification problem (target class vs surrogate negatives)
and samples new points by following energy gradients.

```python theme={null}
TabEBM.generate(
    X: np.ndarray | torch.Tensor | pd.DataFrame,
    y: np.ndarray | torch.Tensor | pd.Series,
    num_samples: int,
    starting_point_noise_std: float = 0.01,
    sgld_step_size: float = 0.1,
    sgld_noise_std: float = 0.01,
    sgld_steps: int = 200,
    distance_negative_class: float = 5,
    seed: int = 42,
    verbose: bool = False,
) -> dict[str, np.ndarray]
```

**Parameters**

<div className="python-reference-table">
  | Parameter | Type | Default | Description |
  | - | - | - | - |
  | <span id="tabebm-tabebm-tabebm-generate--x" /><code className="python-reference-parameter">X</code> | <code className="python-reference-type"><a href="https://numpy.org/doc/stable/reference/generated/numpy.ndarray.html">np.ndarray</a> \| torch.Tensor \| <a href="https://pandas.pydata.org/docs/reference/api/pandas.DataFrame.html">pd.Data<wbr />Frame</a></code> | Required | Input features of shape (n\_samples, n\_features) |
  | <span id="tabebm-tabebm-tabebm-generate--y" /><code className="python-reference-parameter">y</code> | <code className="python-reference-type"><a href="https://numpy.org/doc/stable/reference/generated/numpy.ndarray.html">np.ndarray</a> \| torch.Tensor \| <a href="https://pandas.pydata.org/docs/reference/api/pandas.Series.html">pd.Series</a></code> | Required | Target labels of shape (n\_samples,) |
  | <span id="tabebm-tabebm-tabebm-generate--num-samples" /><code className="python-reference-parameter">num\_<wbr />samples</code> | <code className="python-reference-type">int</code> | Required | Number of synthetic samples to generate per class |
  | <span id="tabebm-tabebm-tabebm-generate--starting-point-noise-std" /><code className="python-reference-parameter">starting\_<wbr />point\_<wbr />noise\_<wbr />std</code> | <code className="python-reference-type">float</code> | `0.01` | Standard deviation of noise added to real data                     points when initializing SGLD chains |
  | <span id="tabebm-tabebm-tabebm-generate--sgld-step-size" /><code className="python-reference-parameter">sgld\_<wbr />step\_<wbr />size</code> | <code className="python-reference-type">float</code> | `0.1` | Step size for gradient updates in SGLD sampling |
  | <span id="tabebm-tabebm-tabebm-generate--sgld-noise-std" /><code className="python-reference-parameter">sgld\_<wbr />noise\_<wbr />std</code> | <code className="python-reference-type">float</code> | `0.01` | Standard deviation of noise added at each SGLD step |
  | <span id="tabebm-tabebm-tabebm-generate--sgld-steps" /><code className="python-reference-parameter">sgld\_<wbr />steps</code> | <code className="python-reference-type">int</code> | `200` | Number of SGLD steps to perform |
  | <span id="tabebm-tabebm-tabebm-generate--distance-negative-class" /><code className="python-reference-parameter">distance\_<wbr />negative\_<wbr />class</code> | <code className="python-reference-type">float</code> | `5` | Distance for placing surrogate negative samples                    from the origin (used to create binary classification) |
  | <span id="tabebm-tabebm-tabebm-generate--seed" /><code className="python-reference-parameter">seed</code> | <code className="python-reference-type">int</code> | `42` | Random seed for reproducibility |
  | <span id="tabebm-tabebm-tabebm-generate--verbose" /><code className="python-reference-parameter">verbose</code> | <code className="python-reference-type">bool</code> | `False` | Whether to print verbose information during sampling |
</div>

**Returns**

<div className="python-reference-table python-reference-returns">
  | Type | Description |
  | - | - |
  | <code className="python-reference-type">dict\[str, <a href="https://numpy.org/doc/stable/reference/generated/numpy.ndarray.html">np.ndarray</a>]</code> | Dictionary mapping class names to generated samples: \{     'class\_0': [`np.ndarray`](https://numpy.org/doc/stable/reference/generated/numpy.ndarray.html) of shape (`num_samples`, n\_features),     'class\_1': [`np.ndarray`](https://numpy.org/doc/stable/reference/generated/numpy.ndarray.html) of shape (`num_samples`, n\_features),     ... } |
</div>

***

<div className="python-reference-heading">
  <h2 id="tabebm-tabebm-tabebm-train-test-split-allow-full-train">
    `TabEBM.train_test_split_allow_full_train`
  </h2>

  <a className="python-reference-source" href="https://github.com/PriorLabs/tabpfn-extensions/blob/840c15a1848a986b39c85bc17efc61e0e377f983/src/tabpfn_extensions/tabebm/tabebm.py#L669" aria-label="View source for TabEBM.train_test_split_allow_full_train"><span aria-hidden="true">\</></span> View source <span aria-hidden="true">↗</span></a>
</div>

Enhanced train-test split that supports full training mode.

This method extends sklearn's train\_test\_split to handle the case where
`test_size`=0, which means we want to use all data for training (no validation).
This is useful for [`TabEBM`](/api-reference/python/tabpfn-extensions/tabebm#tabebm-tabebm-tabebm)'s energy-based training approach.

```python theme={null}
TabEBM.train_test_split_allow_full_train(
    X: np.ndarray | torch.Tensor,
    y: np.ndarray | torch.Tensor,
    test_size: float | None = None,
    train_size: float | None = None,
    random_state: int | None = None,
    shuffle: bool = True,
    stratify: np.ndarray | torch.Tensor | None = None,
) -> tuple[np.ndarray | torch.Tensor, np.ndarray | torch.Tensor, np.ndarray | torch.Tensor, np.ndarray | torch.Tensor]
```

**Parameters**

<div className="python-reference-table">
  | Parameter | Type | Default | Description |
  | - | - | - | - |
  | <span id="tabebm-tabebm-tabebm-train-test-split-allow-full-train--x" /><code className="python-reference-parameter">X</code> | <code className="python-reference-type"><a href="https://numpy.org/doc/stable/reference/generated/numpy.ndarray.html">np.ndarray</a> \| torch.Tensor</code> | Required | — |
  | <span id="tabebm-tabebm-tabebm-train-test-split-allow-full-train--y" /><code className="python-reference-parameter">y</code> | <code className="python-reference-type"><a href="https://numpy.org/doc/stable/reference/generated/numpy.ndarray.html">np.ndarray</a> \| torch.Tensor</code> | Required | — |
  | <span id="tabebm-tabebm-tabebm-train-test-split-allow-full-train--test-size" /><code className="python-reference-parameter">test\_<wbr />size</code> | <code className="python-reference-type">float \| None</code> | `None` | Fraction of data for testing (if 0, enables full train mode) |
  | <span id="tabebm-tabebm-tabebm-train-test-split-allow-full-train--train-size" /><code className="python-reference-parameter">train\_<wbr />size</code> | <code className="python-reference-type">float \| None</code> | `None` | Fraction of data for training |
  | <span id="tabebm-tabebm-tabebm-train-test-split-allow-full-train--random-state" /><code className="python-reference-parameter">random\_<wbr />state</code> | <code className="python-reference-type">int \| None</code> | `None` | Random seed for reproducibility |
  | <span id="tabebm-tabebm-tabebm-train-test-split-allow-full-train--shuffle" /><code className="python-reference-parameter">shuffle</code> | <code className="python-reference-type">bool</code> | `True` | Whether to shuffle data before splitting |
  | <span id="tabebm-tabebm-tabebm-train-test-split-allow-full-train--stratify" /><code className="python-reference-parameter">stratify</code> | <code className="python-reference-type"><a href="https://numpy.org/doc/stable/reference/generated/numpy.ndarray.html">np.ndarray</a> \| torch.Tensor \| None</code> | `None` | Array-like for stratified splitting |
</div>

**Returns**

<div className="python-reference-table python-reference-returns">
  | Type | Description |
  | - | - |
  | <code className="python-reference-type">tuple\[<a href="https://numpy.org/doc/stable/reference/generated/numpy.ndarray.html">np.ndarray</a> \| torch.Tensor, <a href="https://numpy.org/doc/stable/reference/generated/numpy.ndarray.html">np.ndarray</a> \| torch.Tensor, <a href="https://numpy.org/doc/stable/reference/generated/numpy.ndarray.html">np.ndarray</a> \| torch.Tensor, <a href="https://numpy.org/doc/stable/reference/generated/numpy.ndarray.html">np.ndarray</a> \| torch.Tensor]</code> | Tuple of (X\_train, X\_val, y\_train, y\_val) In full train mode, X\_train=`X` and y\_train=y |
</div>

***

<div className="python-reference-heading">
  <h2 id="tabebm-tabebm-seed-everything">
    `seed_everything`
  </h2>

  <a className="python-reference-source" href="https://github.com/PriorLabs/tabpfn-extensions/blob/840c15a1848a986b39c85bc17efc61e0e377f983/src/tabpfn_extensions/tabebm/tabebm.py#L45" aria-label="View source for seed_everything"><span aria-hidden="true">\</></span> View source <span aria-hidden="true">↗</span></a>
</div>

Set random seeds for reproducibility across all libraries.

```python theme={null}
seed_everything(
    seed: int,
) -> None
```

**Parameters**

<div className="python-reference-table">
  | Parameter | Type | Default | Description |
  | - | - | - | - |
  | <span id="tabebm-tabebm-seed-everything--seed" /><code className="python-reference-parameter">seed</code> | <code className="python-reference-type">int</code> | Required | Random seed value |
</div>

**Returns**

<div className="python-reference-table python-reference-returns">
  | Type | Description |
  | - | - |
  | <code className="python-reference-type">None</code> | — |
</div>

***

<div className="python-reference-heading">
  <h2 id="tabebm-tabebm-to-numpy">
    `to_numpy`
  </h2>

  <a className="python-reference-source" href="https://github.com/PriorLabs/tabpfn-extensions/blob/840c15a1848a986b39c85bc17efc61e0e377f983/src/tabpfn_extensions/tabebm/tabebm.py#L21" aria-label="View source for to_numpy"><span aria-hidden="true">\</></span> View source <span aria-hidden="true">↗</span></a>
</div>

Convert input data to numpy array format.

```python theme={null}
to_numpy(
    X: np.ndarray | torch.Tensor | pd.DataFrame | None,
) -> np.ndarray | None
```

**Parameters**

<div className="python-reference-table">
  | Parameter | Type | Default | Description |
  | - | - | - | - |
  | <span id="tabebm-tabebm-to-numpy--x" /><code className="python-reference-parameter">X</code> | <code className="python-reference-type"><a href="https://numpy.org/doc/stable/reference/generated/numpy.ndarray.html">np.ndarray</a> \| torch.Tensor \| <a href="https://pandas.pydata.org/docs/reference/api/pandas.DataFrame.html">pd.Data<wbr />Frame</a> \| None</code> | Required | Input data in various formats (numpy array, torch tensor, pandas DataFrame, or `None`) |
</div>

**Returns**

<div className="python-reference-table python-reference-returns">
  | Type | Description |
  | - | - |
  | <code className="python-reference-type"><a href="https://numpy.org/doc/stable/reference/generated/numpy.ndarray.html">np.ndarray</a> \| None</code> | numpy array representation of the input data, or `None` if input is `None` |
</div>

**Raises**

`ValueError`

If input type is not supported


This documentation is built and hosted on [Mintlify](https://mintlify.com), a developer documentation platform.