Skip to content

API reference

Public objects are imported from locata_torch. The reference below is generated from the package source by mkdocstrings; private parsing helpers are excluded.

Read getting started for constructor defaults and selection, data model for shapes and missing values, time and windows for clock semantics, and I/O and DataLoader for caching, collation, and worker behavior.

locata_torch

Read LOCATA using map-style PyTorch datasets and explicit clock contracts.

CalendarTime module-attribute

CalendarTime = tuple[int, int, int, int, int, float]

LocataDataset

LocataDataset(
    root: str | Path | None = None,
    *,
    split: str | Sequence[str] = "dev",
    tasks: Sequence[int] = (1, 2, 3, 4, 5, 6),
    recordings: Sequence[int] | None = None,
    arrays: Sequence[str] | None = None,
    dtype: dtype = torch.float32,
    load_source_audio: bool = False,
    cache_size: int = 16,
    cache_bytes: int = 16777216,
)

Bases: Dataset[LocataSample]

One item per existing (split, task, recording, array) WAV.

Payloads are read lazily. Source audio is opt-in; geometry and available VAD are returned independently. All returned tensors are on the CPU. See the data model and time-and-windows guides for detailed shapes and clock contracts.

Parameters:

Name Type Description Default
root str | Path | None

Existing LOCATA root, or None to use LOCATA_ROOT then managed storage. User-home expansion is supported. Never triggers a download.

None
split str | Sequence[str]

One split name, or a sequence selecting dev and/or eval.

'dev'
tasks Sequence[int]

Task numbers from 1 through 6.

(1, 2, 3, 4, 5, 6)
recordings Sequence[int] | None

Positive recording numbers across selected tasks, or all.

None
arrays Sequence[str] | None

Supported array names, or all available arrays.

None
dtype dtype

Waveform dtype, either torch.float32 or torch.float64.

float32
load_source_audio bool

Read available source WAV payloads when true.

False
cache_size int

Maximum cached sparse TXT indexes per worker.

16
cache_bytes int

Maximum estimated cached index bytes per worker.

16777216

Attributes:

Name Type Description
index tuple[RecordingInfo, ...]

Immutable recording information, sorted by split, numeric task, numeric recording, and array name.

missing_audio tuple[Path, ...]

Existing selected array directories without an array WAV.

Raises:

Type Description
LocataError

The root, a WAV header, or a mandatory input is invalid.

ValueError

A selection, dtype, or cache limit is invalid.

Source code in src/locata_torch/dataset.py
def __init__(
    self,
    root: str | Path | None = None,
    *,
    split: str | Sequence[str] = "dev",
    tasks: Sequence[int] = (1, 2, 3, 4, 5, 6),
    recordings: Sequence[int] | None = None,
    arrays: Sequence[str] | None = None,
    dtype: torch.dtype = torch.float32,
    load_source_audio: bool = False,
    cache_size: int = 16,
    cache_bytes: int = 16777216,
):
    splits = select_splits(split)
    self.root = resolve_root(root, splits)
    if not tasks or any(_positive(task, "task") > 6 for task in tasks):
        raise ValueError("tasks must select integers from 1 through 6")
    if recordings is not None and (
        not recordings or any(_positive(rec, "recording") < 1 for rec in recordings)
    ):
        raise ValueError("recordings must select positive integers")
    chosen_arrays = ARRAYS if arrays is None else frozenset(arrays)
    if not chosen_arrays or not chosen_arrays <= ARRAYS:
        raise ValueError(f"arrays must select from {sorted(ARRAYS)}")
    if dtype not in (torch.float32, torch.float64):
        raise ValueError("dtype must be torch.float32 or torch.float64")
    if any(
        not isinstance(v, int) or isinstance(v, bool) or v < 0
        for v in (cache_size, cache_bytes)
    ):
        raise ValueError("cache_size and cache_bytes must be nonnegative integers")
    self.dtype = dtype
    self.load_source_audio = load_source_audio
    self.cache_size = cache_size
    self.cache_bytes = cache_bytes
    self._cache: OrderedDict[Path, SparseTable] = OrderedDict()
    self._cache_pid = os.getpid()
    index: list[RecordingInfo] = []
    missing: list[Path] = []
    for split_name in sorted(set(splits)):
        for task in sorted(set(tasks)):
            task_path = self.root / split_name / f"task{task}"
            if not task_path.is_dir():
                continue
            candidates = [
                (int(match[1]), directory)
                for directory in task_path.iterdir()
                if directory.is_dir()
                and (match := re.fullmatch(r"recording(\d+)", directory.name))
            ]
            for rec, directory in sorted(candidates):
                if recordings is not None and rec not in recordings:
                    continue
                for array in sorted(chosen_arrays):
                    path = directory / array
                    if not path.is_dir():
                        continue
                    audio_path = path / f"audio_array_{array}.wav"
                    if not audio_path.is_file():
                        missing.append(path)
                        continue
                    info = _info(audio_path)
                    for name in (
                        f"audio_array_timestamps_{array}.txt",
                        f"position_array_{array}.txt",
                        "required_time.txt",
                    ):
                        if not (path / name).is_file():
                            raise LocataError(
                                f"{path / name}: required file is missing"
                            )
                    index.append(
                        RecordingInfo(
                            split_name,
                            task,
                            rec,
                            array,
                            path,
                            info.frames,
                            info.samplerate,
                            info.channels,
                        )
                    )
    self.index = tuple(index)
    self.missing_audio = tuple(missing)
    if missing:
        warnings.warn(
            "Array directories without WAV: " + ", ".join(map(str, missing)),
            MissingAudioWarning,
            stacklevel=2,
        )

root instance-attribute

root = resolve_root(root, splits)

dtype instance-attribute

dtype = dtype

load_source_audio instance-attribute

load_source_audio = load_source_audio

cache_size instance-attribute

cache_size = cache_size

cache_bytes instance-attribute

cache_bytes = cache_bytes

index instance-attribute

index = tuple(index)

missing_audio instance-attribute

missing_audio = tuple(missing)

windows

windows(
    *,
    num_samples: int,
    hop_samples: int | None = None,
    drop_last: bool = True,
) -> LocataWindowDataset

Create a partial-read window view with the recording's time origin.

Parameters:

Name Type Description Default
num_samples int

Positive window length in WAV frames.

required
hop_samples int | None

Positive hop in frames, or the window length if omitted.

None
drop_last bool

Keep only complete windows when true; retain short tails without padding when false.

True

Returns:

Type Description
LocataWindowDataset

A map-style view cropping independent annotations to half-open

LocataWindowDataset

window time bounds.

Source code in src/locata_torch/dataset.py
def windows(
    self,
    *,
    num_samples: int,
    hop_samples: int | None = None,
    drop_last: bool = True,
) -> "LocataWindowDataset":
    """Create a partial-read window view with the recording's time origin.

    Args:
        num_samples: Positive window length in WAV frames.
        hop_samples: Positive hop in frames, or the window length if omitted.
        drop_last: Keep only complete windows when true; retain short tails
            without padding when false.

    Returns:
        A map-style view cropping independent annotations to half-open
        window time bounds.
    """
    return LocataWindowDataset(
        self, num_samples=num_samples, hop_samples=hop_samples, drop_last=drop_last
    )

LocataWindowDataset

LocataWindowDataset(
    dataset: LocataDataset,
    *,
    num_samples: int,
    hop_samples: int | None = None,
    drop_last: bool = True,
)

Bases: Dataset[LocataSample]

Fixed-frame view; every waveform read seeks directly to its window.

Prefer constructing this view with LocataDataset.windows.

Parameters:

Name Type Description Default
dataset LocataDataset

Parent recording dataset and its I/O options.

required
num_samples int

Positive window length in WAV frames.

required
hop_samples int | None

Positive frame hop, or the window length if omitted.

None
drop_last bool

Exclude incomplete tails when true. Otherwise return each start inside the recording, with the final frames left unpadded.

True

Raises:

Type Description
ValueError

The window length or hop is not a positive integer.

Source code in src/locata_torch/dataset.py
def __init__(
    self,
    dataset: LocataDataset,
    *,
    num_samples: int,
    hop_samples: int | None = None,
    drop_last: bool = True,
):
    self.dataset = dataset
    self.num_samples = _positive(num_samples, "num_samples")
    self.hop_samples = (
        num_samples
        if hop_samples is None
        else _positive(hop_samples, "hop_samples")
    )
    self.drop_last = drop_last
    self._ends = [0]
    for record in dataset.index:
        n = record.num_frames
        count = (
            max(0, (n - num_samples) // self.hop_samples + 1)
            if drop_last
            else (n - 1) // self.hop_samples + 1
        )
        self._ends.append(self._ends[-1] + count)

dataset instance-attribute

dataset = dataset

num_samples instance-attribute

num_samples = _positive(num_samples, 'num_samples')

hop_samples instance-attribute

hop_samples = (
    num_samples
    if hop_samples is None
    else _positive(hop_samples, "hop_samples")
)

drop_last instance-attribute

drop_last = drop_last

LocataError

Bases: RuntimeError

Invalid LOCATA input, configuration, or preparation, with its path and cause.

MissingAudioWarning

Bases: UserWarning

An existing array directory has no matching array WAV.

RecordingInfo dataclass

RecordingInfo(
    split: str,
    task: int,
    recording: int,
    array: str,
    path: Path,
    num_frames: int,
    sample_rate: int,
    num_channels: int,
)

One existing array WAV, indexed without loading its payload.

split instance-attribute

split: str

task instance-attribute

task: int

recording instance-attribute

recording: int

array instance-attribute

array: str

path instance-attribute

path: Path

num_frames instance-attribute

num_frames: int

sample_rate instance-attribute

sample_rate: int

num_channels instance-attribute

num_channels: int

id property

id: str

audio_path property

audio_path: Path

RecordingMetadata

Bases: TypedDict

Recording identity and half-open frame and relative-time boundaries.

id instance-attribute

id: str

split instance-attribute

split: str

task instance-attribute

task: int

recording instance-attribute

recording: int

array instance-attribute

array: str

path instance-attribute

path: Path

start_frame instance-attribute

start_frame: int

stop_frame instance-attribute

stop_frame: int

num_frames instance-attribute

num_frames: int

time_bounds instance-attribute

time_bounds: tuple[float, float]

LocataSample

Bases: TypedDict

One recording or window, preserving independent annotation clocks.

waveform instance-attribute

waveform: Tensor

sample_rate instance-attribute

sample_rate: int

audio_time instance-attribute

audio_time: Tensor

time_origin instance-attribute

time_origin: CalendarTime

required_time instance-attribute

required_time: RequiredTime

array_pose instance-attribute

array_pose: ArrayPose

sources instance-attribute

sources: dict[str, Source]

metadata instance-attribute

metadata: RecordingMetadata

LocataBatch

Bases: TypedDict

Padded audio and masks with original annotation lists per sample.

waveform instance-attribute

waveform: Tensor

lengths instance-attribute

lengths: Tensor

audio_mask instance-attribute

audio_mask: Tensor

sample_rate instance-attribute

sample_rate: int

audio_time instance-attribute

audio_time: list[Tensor]

time_origin instance-attribute

time_origin: list[CalendarTime]

required_time instance-attribute

required_time: list[RequiredTime]

array_pose instance-attribute

array_pose: list[ArrayPose]

sources instance-attribute

sources: list[dict[str, Source]]

metadata instance-attribute

metadata: list[RecordingMetadata]

RequiredTime

Bases: TypedDict

Requested estimation times and validity flags, independent of VAD.

time instance-attribute

time: Tensor

valid_flag instance-attribute

valid_flag: Tensor

Pose

Bases: TypedDict

Timed world pose: xyz, reference vector, and local-to-world rotation.

time instance-attribute

time: Tensor

time_origin instance-attribute

time_origin: CalendarTime

position instance-attribute

position: Tensor

ref_vec instance-attribute

ref_vec: Tensor

rotation instance-attribute

rotation: Tensor

ArrayPose

Bases: Pose

Array pose with microphone world positions in WAV channel order.

microphone_position instance-attribute

microphone_position: Tensor

Source

Bases: TypedDict

Independent source pose, opt-in audio, availability, and optional VAD.

pose instance-attribute

pose: Pose | None

audio instance-attribute

audio: SourceAudio | None

audio_available instance-attribute

audio_available: bool

vad instance-attribute

vad: SourceVAD

SourceAudio

Bases: TypedDict

Optional source waveform and its independent, shared-origin clock.

waveform instance-attribute

waveform: Tensor

sample_rate instance-attribute

sample_rate: int

audio_time instance-attribute

audio_time: Tensor

start_frame instance-attribute

start_frame: int

stop_frame instance-attribute

stop_frame: int

SourceVAD

Bases: TypedDict

Optional array-aligned and source-aligned VAD for one source ID.

array instance-attribute

array: TimedVAD | None

source instance-attribute

source: TimedVAD | None

TimedVAD

Bases: TypedDict

Known boolean voice activity and its corresponding audio timestamps.

time instance-attribute

time: Tensor

values instance-attribute

values: Tensor

DOA

Bases: TypedDict

Array-frame vector, LOCATA azimuth, inclination from +z, and range.

vector instance-attribute

vector: Tensor

azimuth instance-attribute

azimuth: Tensor

inclination instance-attribute

inclination: Tensor

range instance-attribute

range: Tensor

download_locata

download_locata(
    split: str | Sequence[str] = "dev",
    *,
    data_dir: str | Path | None = None,
) -> Path

Install dev/eval from the pinned official release and return its root.

Storage lookup is data_dir, then LOCATA_DATA_DIR, then the platform user-data directory. LOCATA_ROOT does not affect downloads. Each split is verified and installed independently. Existing managed data is checked and reused; unknown or altered data raises :class:LocataError.

Requires locata-torch[download]. This explicit operation performs network I/O; Dataset construction never downloads. Call it before starting workers.

Source code in src/locata_torch/_download.py
def download_locata(
    split: str | Sequence[str] = "dev", *, data_dir: str | Path | None = None
) -> Path:
    """Install dev/eval from the pinned official release and return its root.

    Storage lookup is ``data_dir``, then ``LOCATA_DATA_DIR``, then the platform
    user-data directory. ``LOCATA_ROOT`` does not affect downloads. Each split
    is verified and installed independently. Existing managed data is checked
    and reused; unknown or altered data raises :class:`LocataError`.

    Requires ``locata-torch[download]``. This explicit operation performs network
    I/O; Dataset construction never downloads. Call it before starting workers.
    """
    splits = select_splits(split)
    try:
        from filelock import FileLock, Timeout
    except ImportError as exc:
        raise LocataError(
            'install "locata-torch[download]" to enable downloading'
        ) from exc
    store = data_directory(data_dir)
    try:
        _check_store(store)
        required = 0
        for value in splits:
            manifest = read_manifest(store, value)
            target = dataset_directory(store) / value
            if manifest is None:
                if target.exists() or target.is_symlink():
                    raise LocataError(
                        f"{target}: unknown existing split without a manifest"
                    )
                spec = ARCHIVES[value]
                cached = store / "archives" / RELEASE / f"{value}.zip"
                retained = (
                    spec.size
                    if cached.is_file()
                    and not cached.is_symlink()
                    and cached.stat().st_size == spec.size
                    else 0
                )
                required += spec.size - retained + spec.expanded_size
        if required:
            _space(store, required)
        for relative in (
            f"archives/{RELEASE}",
            f"datasets/{RELEASE}",
            f"state/{RELEASE}",
            "staging",
            "locks",
        ):
            (store / relative).mkdir(parents=True, exist_ok=True)
        for value in splits:
            lock_path = store / "locks" / f"{RELEASE}-{value}.lock"
            if lock_path.is_symlink():
                raise LocataError(f"{lock_path}: lock file is a symlink")
            with FileLock(lock_path, timeout=60):
                _check_store(store)
                manifest = read_manifest(store, value)
                if manifest is not None:
                    _recover(store, value, manifest)
                else:
                    target = dataset_directory(store) / value
                    if target.exists() or target.is_symlink():
                        raise LocataError(f"{target}: unknown existing split")
                    _install(store, ARCHIVES[value])
        return dataset_directory(store)
    except (OSError, ValueError, zipfile.BadZipFile, Timeout) as exc:
        raise LocataError(f"{store}: LOCATA installation failed: {exc}") from exc

collate_locata

collate_locata(
    samples: Sequence[LocataSample],
) -> LocataBatch

Right-pad audio; retain annotations and optional sources per sample.

Parameters:

Name Type Description Default
samples Sequence[LocataSample]

Nonempty samples sharing array, sample rate, channels, and dtype.

required

Returns:

Type Description
LocataBatch

CPU audio [B, C, T_max], int64 lengths, a boolean padding-only mask,

LocataBatch

and lists of unmodified clocks, annotations, sources, and metadata.

Raises:

Type Description
ValueError

Samples are empty, heterogeneous, or have non-CPU waveforms.

Source code in src/locata_torch/collate.py
def collate_locata(samples: Sequence[LocataSample]) -> LocataBatch:
    """Right-pad audio; retain annotations and optional sources per sample.

    Args:
        samples: Nonempty samples sharing array, sample rate, channels, and dtype.

    Returns:
        CPU audio `[B, C, T_max]`, int64 lengths, a boolean padding-only mask,
        and lists of unmodified clocks, annotations, sources, and metadata.

    Raises:
        ValueError: Samples are empty, heterogeneous, or have non-CPU waveforms.
    """
    if not samples:
        raise ValueError("cannot collate an empty batch")
    first = samples[0]
    waveform = first["waveform"]
    signature = (
        first["metadata"]["array"],
        first["sample_rate"],
        waveform.shape[0],
        waveform.dtype,
    )
    for sample in samples:
        audio = sample["waveform"]
        if audio.ndim != 2 or audio.device.type != "cpu":
            raise ValueError("waveform must be a CPU tensor [channels, samples]")
        if (
            sample["metadata"]["array"],
            sample["sample_rate"],
            audio.shape[0],
            audio.dtype,
        ) != signature:
            raise ValueError(
                "batch must share array, sample rate, channel count and dtype"
            )
    lengths = torch.tensor([s["waveform"].shape[1] for s in samples], dtype=torch.int64)
    maximum = int(lengths.max())
    padded = waveform.new_zeros(len(samples), waveform.shape[0], maximum)
    for row, sample in enumerate(samples):
        padded[row, :, : sample["waveform"].shape[1]] = sample["waveform"]
    return {
        "waveform": padded,
        "lengths": lengths,
        "audio_mask": torch.arange(maximum).unsqueeze(0) < lengths.unsqueeze(1),
        "sample_rate": first["sample_rate"],
        "audio_time": [s["audio_time"] for s in samples],
        "time_origin": [s["time_origin"] for s in samples],
        "required_time": [s["required_time"] for s in samples],
        "array_pose": [s["array_pose"] for s in samples],
        "sources": [s["sources"] for s in samples],
        "metadata": [s["metadata"] for s in samples],
    }

world_to_array

world_to_array(
    source_position: Tensor,
    array_position: Tensor,
    array_rotation: Tensor,
) -> Tensor

Apply R.T @ (h - p) to simultaneous, broadcastable Cartesian positions.

Parameters:

Name Type Description Default
source_position Tensor

World xyz in metres, with trailing shape [3].

required
array_position Tensor

Array origin in world metres, with trailing shape [3].

required
array_rotation Tensor

Local-to-world rotation with trailing shape [3, 3].

required

Returns:

Type Description
Tensor

Float64 array-frame vectors with broadcast leading dimensions.

Raises:

Type Description
ValueError

Shapes, finiteness, or proper-rotation checks fail.

Source code in src/locata_torch/geometry.py
def world_to_array(
    source_position: Tensor, array_position: Tensor, array_rotation: Tensor
) -> Tensor:
    """Apply R.T @ (h - p) to simultaneous, broadcastable Cartesian positions.

    Args:
        source_position: World xyz in metres, with trailing shape `[3]`.
        array_position: Array origin in world metres, with trailing shape `[3]`.
        array_rotation: Local-to-world rotation with trailing shape `[3, 3]`.

    Returns:
        Float64 array-frame vectors with broadcast leading dimensions.

    Raises:
        ValueError: Shapes, finiteness, or proper-rotation checks fail.
    """
    if (
        source_position.shape[-1:] != (3,)
        or array_position.shape[-1:] != (3,)
        or array_rotation.shape[-2:] != (3, 3)
    ):
        raise ValueError("position shape must end in 3 and rotation shape in (3, 3)")
    h, p, rotation = (
        tensor.to(dtype=torch.float64)
        for tensor in (source_position, array_position, array_rotation)
    )
    if not all(bool(torch.isfinite(tensor).all()) for tensor in (h, p, rotation)):
        raise ValueError("geometry must be finite")
    identity = torch.eye(3, dtype=torch.float64, device=rotation.device)
    if not torch.allclose(
        rotation.transpose(-1, -2) @ rotation,
        identity.expand_as(rotation),
        atol=1e-5,
        rtol=1e-5,
    ) or not torch.allclose(
        torch.linalg.det(rotation),
        torch.ones_like(torch.linalg.det(rotation)),
        atol=1e-5,
        rtol=1e-5,
    ):
        raise ValueError("rotation must be an orthonormal matrix with determinant +1")
    try:
        return (rotation.transpose(-1, -2) @ (h - p).unsqueeze(-1)).squeeze(-1)
    except RuntimeError as exc:
        raise ValueError(
            "geometry shapes must be broadcastable on the same device"
        ) from exc

locata_doa

locata_doa(array_pose: Pose, source_pose: Pose) -> DOA

Derive angles only for poses sharing exactly the same origin and times.

Parameters:

Name Type Description Default
array_pose Pose

Array world position and local-to-world rotation.

required
source_pose Pose

Source world position at exactly matching timestamps.

required

Returns:

Type Description
DOA

Float64 vectors in metres, azimuth in [-pi, pi) measured from +y,

DOA

inclination in [0, pi] measured from +z, and range in metres.

DOA

Angles use radians; no temporal interpolation is performed.

Raises:

Type Description
ValueError

Clocks or shapes differ, geometry is invalid, or distance is zero.

Source code in src/locata_torch/geometry.py
def locata_doa(array_pose: Pose, source_pose: Pose) -> DOA:
    """Derive angles only for poses sharing exactly the same origin and times.

    Args:
        array_pose: Array world position and local-to-world rotation.
        source_pose: Source world position at exactly matching timestamps.

    Returns:
        Float64 vectors in metres, azimuth in `[-pi, pi)` measured from +y,
        inclination in `[0, pi]` measured from +z, and range in metres.
        Angles use radians; no temporal interpolation is performed.

    Raises:
        ValueError: Clocks or shapes differ, geometry is invalid, or distance
            is zero.
    """
    if array_pose["time_origin"] != source_pose["time_origin"]:
        raise ValueError("pose time origins must match")
    if not torch.equal(array_pose["time"], source_pose["time"]):
        raise ValueError("pose times must match exactly; align them explicitly")
    rows = len(array_pose["time"])
    if (
        array_pose["position"].shape != (rows, 3)
        or source_pose["position"].shape != (rows, 3)
        or array_pose["rotation"].shape != (rows, 3, 3)
    ):
        raise ValueError("pose shape must match its time rows")
    vector = world_to_array(
        source_pose["position"], array_pose["position"], array_pose["rotation"]
    )
    distance = torch.linalg.vector_norm(vector, dim=-1)
    if (distance == 0).any():
        raise ValueError("DOA is undefined at zero distance")
    azimuth = torch.atan2(vector[:, 1], vector[:, 0]) - math.pi / 2
    return {
        "vector": vector,
        "azimuth": torch.remainder(azimuth + math.pi, 2 * math.pi) - math.pi,
        "inclination": torch.acos((vector[:, 2] / distance).clamp(-1, 1)),
        "range": distance,
    }