diff --git a/python/fasttext_module/fasttext/FastText.py b/python/fasttext_module/fasttext/FastText.py index d4081b4..f66ada1 100644 --- a/python/fasttext_module/fasttext/FastText.py +++ b/python/fasttext_module/fasttext/FastText.py @@ -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 diff --git a/python/fasttext_module/fasttext/tests/test_autotune_bucket.py b/python/fasttext_module/fasttext/tests/test_autotune_bucket.py new file mode 100644 index 0000000..44f06cd --- /dev/null +++ b/python/fasttext_module/fasttext/tests/test_autotune_bucket.py @@ -0,0 +1,68 @@ +# 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 +import fasttext_pybind +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) + + +def test_unsupervised_autotune_with_manual_bucket_zero(tmp_path): + """Used to terminate: trial 1 kept the default subwords (CLI-only path).""" + path = tmp_path / "train.txt" + path.write_text("".join(f"{line}\n" for line in get_random_data(3000))) + # train_unsupervised() takes no autotune args; the CLI reaches this. + args = dict( + FastText.unsupervised_default, + input=str(path), + autotuneValidationFile=str(path), + autotuneDuration=3, + bucket=0, + lr=0.05, # a sampled lr can diverge (NaN) in the final retrain + thread=12, # thread <= 10 leaves the input matrix partly uninitialized + verbose=0, + ) + a = FastText._build_args(args, {"bucket", "lr"}) + model = FastText._FastText(args=a) + fasttext_pybind.train(model.f, a) + got = model.f.getArgs() + assert (got.bucket, got.wordNgrams, got.maxn) == (0, 1, 0) diff --git a/src/autotune.cc b/src/autotune.cc index 567731b..ad1dc50 100644 --- a/src/autotune.cc +++ b/src/autotune.cc @@ -120,7 +120,17 @@ AutotuneStrategy::AutotuneStrategy( bestNonzeroBucket_(2000000), originalBucket_(originalArgs.bucket) { minnChoices_ = {0, 2, 3}; - updateBest(originalArgs); + Args args = originalArgs; + // A manual bucket=0 means no n-gram buckets: start without n-grams too. + if (args.isManual("bucket") && args.bucket == 0) { + if (!args.isManual("wordNgrams")) { + args.wordNgrams = 1; + } + if (!args.isManual("maxn")) { + args.maxn = 0; + } + } + updateBest(args); } Args AutotuneStrategy::ask(double elapsed) { @@ -132,6 +142,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_); @@ -142,7 +154,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_); } @@ -151,7 +163,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, @@ -164,7 +176,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;