Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 7 additions & 1 deletion python/fasttext_module/fasttext/FastText.py
Original file line number Diff line number Diff line change
Expand Up @@ -429,7 +429,13 @@ def _build_args(args, manually_set_args):
a.setManual(k)
a.output = "" # User should use save_model
a.saveOutput = 0 # Never use this
if a.wordNgrams <= 1 and a.maxn == 0:
# Autotune reuses a manual bucket for the n-gram trials it samples.
keep_bucket = (
bool(a.autotuneValidationFile)
and "bucket" in manually_set_args
and a.bucket > 0
)
if a.wordNgrams <= 1 and a.maxn == 0 and not keep_bucket:
a.bucket = 0
return a

Expand Down
45 changes: 45 additions & 0 deletions python/fasttext_module/fasttext/tests/test_autotune_bucket.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,45 @@
# SPDX-FileContributor: Arthit Suriyawongkul
# SPDX-FileCopyrightText: 2026-present, fasttext-community
# SPDX-FileType: SOURCE
# SPDX-License-Identifier: MIT

"""Autotune must respect a manually set bucket."""

import fasttext
from fasttext import FastText

from .helpers import get_random_data


def test_build_args_keeps_manual_bucket_for_autotune():
"""Used to zero it, so autotune's n-gram trials hashed modulo 0."""
# unsupervised_default lists every arg; the overrides make it supervised.
base = dict(
FastText.unsupervised_default,
input="x",
model="supervised",
minn=0,
maxn=0,
autotuneValidationFile="valid.txt",
)
assert FastText._build_args(dict(base, bucket=1000), {"bucket"}).bucket == 1000
assert FastText._build_args(dict(base), set()).bucket == 0 # default: unchanged


def test_autotune_with_manual_bucket_zero(tmp_path):
"""Used to terminate: trials sampled n-grams that bucket=0 can't hold."""
path = tmp_path / "train.txt"
path.write_text("".join(f"__label__{line}\n" for line in get_random_data(3000)))
# thread=12: thread <= 10 leaves the input matrix partly uninitialized.
model = fasttext.train_supervised(
str(path),
autotuneValidationFile=str(path),
autotuneDuration=3, # enough for several trials
bucket=0,
minn=2, # a manual minn must not turn on subwords either
lr=0.1, # a sampled lr can diverge (NaN) in the final retrain
thread=12,
verbose=0,
)
args = model.f.getArgs()
assert (args.bucket, args.wordNgrams, args.maxn) == (0, 1, 0)
8 changes: 5 additions & 3 deletions src/autotune.cc
Original file line number Diff line number Diff line change
Expand Up @@ -132,6 +132,8 @@ Args AutotuneStrategy::ask(double elapsed) {
}

Args args = bestArgs_;
// A manual bucket=0 means no n-gram buckets: don't sample n-grams.
const bool noBuckets = args.isManual("bucket") && originalBucket_ == 0;

if (!args.isManual("epoch")) {
args.epoch = updateArgGauss(args.epoch, 1, 100, 2.8, 2.5, t, false, rng_);
Expand All @@ -142,7 +144,7 @@ Args AutotuneStrategy::ask(double elapsed) {
if (!args.isManual("dim")) {
args.dim = updateArgGauss(args.dim, 1, 1000, 1.4, 0.3, t, false, rng_);
}
if (!args.isManual("wordNgrams")) {
if (!args.isManual("wordNgrams") && !noBuckets) {
args.wordNgrams =
updateArgGauss(args.wordNgrams, 1, 5, 4.3, 2.4, t, true, rng_);
}
Expand All @@ -151,7 +153,7 @@ Args AutotuneStrategy::ask(double elapsed) {
updateArgGauss(bestDsubExponent_, 1, 4, 2.0, 1.0, t, true, rng_);
args.dsub = (1 << dsubExponent);
}
if (!args.isManual("minn")) {
if (!args.isManual("minn") && !noBuckets) {
int minnIndex = updateArgGauss(
bestMinnIndex_,
0,
Expand All @@ -164,7 +166,7 @@ Args AutotuneStrategy::ask(double elapsed) {
args.minn = minnChoices_[minnIndex];
}
if (!args.isManual("maxn")) {
if (args.minn == 0) {
if (args.minn == 0 || noBuckets) {
args.maxn = 0;
} else {
args.maxn = args.minn + 3;
Expand Down
Loading