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
1 change: 1 addition & 0 deletions .gitattributes
Original file line number Diff line number Diff line change
@@ -1 +1,2 @@
*.ipynb eol=lf
*.sh eol=lf
2 changes: 1 addition & 1 deletion local_check.sh
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@ if [ "$#" -ne 0 ]; then
exit 2
fi

poetry install
poetry install --all-extras

if [ "$agent_strict" = true ]; then
echo "=================== comment hygiene ================="
Expand Down
12 changes: 9 additions & 3 deletions machine/jobs/huggingface/hugging_face_nmt_model_factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,8 +6,8 @@
import datasets.utils.logging as datasets_logging
import transformers.utils.logging as transformers_logging
from transformers import AutoConfig, AutoModelForSeq2SeqLM, HfArgumentParser, PreTrainedModel, Seq2SeqTrainingArguments
from transformers.integrations import ClearMLCallback
from transformers.tokenization_utils import TruncationStrategy
from transformers.integrations.integration_utils import ClearMLCallback
from transformers.tokenization_utils_base import TruncationStrategy

from ...corpora.parallel_text_corpus import ParallelTextCorpus
from ...corpora.text_corpus import TextCorpus
Expand All @@ -26,7 +26,13 @@ def __init__(self, config: Any) -> None:
self._config = config
args = config.huggingface.train_params.to_dict()
args["output_dir"] = str(self._model_dir)
args["overwrite_output_dir"] = True
# Allow group_by_length backwards compatibility. The settings default for train_sampling_strategy is
# group_by_length, so any other value was set explicitly and takes precedence over the legacy option.
group_by_length = args.pop("group_by_length", None)
if group_by_length is not None:
logger.warning("'group_by_length' is deprecated. Use 'train_sampling_strategy' instead.")
if args.get("train_sampling_strategy", "group_by_length") == "group_by_length":
args["train_sampling_strategy"] = "group_by_length" if group_by_length else "random"
# Use "max_steps" from root for backward compatibility
if "max_steps" in self._config.huggingface:
args["max_steps"] = self._config.huggingface.max_steps
Expand Down
1 change: 1 addition & 0 deletions machine/jobs/nmt_build_options.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@ class TrainParams(BaseModel):
gradient_accumulation_steps: int | None = None
label_smoothing_factor: float | None = None
group_by_length: bool | None = None
train_sampling_strategy: str | None = None
Comment thread
pmachapman marked this conversation as resolved.
gradient_checkpointing: bool | None = None
lr_scheduler_type: str | None = None
learning_rate: float | None = None
Expand Down
3 changes: 2 additions & 1 deletion machine/jobs/nmt_engine_build_job.py
Original file line number Diff line number Diff line change
Expand Up @@ -149,7 +149,8 @@ def _translate(
if check_canceled is not None:
check_canceled()
source_segments = [pt_info["translation"] for pt_info in pt_batch]
for pt_info, result in zip(pt_batch, engine.translate_batch(source_segments), strict=True):
t_batch = engine.translate_batch(source_segments)
for pt_info, result in zip(pt_batch, t_batch, strict=True):
pt_info["translation"] = result.translation
pt_info["sequenceConfidence"] = result.sequence_confidence
current_inference_step += len(pt_batch)
Expand Down
2 changes: 1 addition & 1 deletion machine/jobs/settings.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@ default:
per_device_train_batch_size: 64
gradient_accumulation_steps: 1
label_smoothing_factor: 0.2
group_by_length: true
train_sampling_strategy: group_by_length
gradient_checkpointing: true
lr_scheduler_type: cosine
learning_rate: 0.0002
Expand Down
10 changes: 8 additions & 2 deletions machine/translation/huggingface/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,8 +6,14 @@
if not is_torch_available():
raise RuntimeError("torch is not installed.")

from .hugging_face_nmt_engine import HuggingFaceNmtEngine
from .hugging_face_nmt_engine import HuggingFaceNmtEngine, SilTranslationPipeline
from .hugging_face_nmt_model import HuggingFaceNmtModel
from .hugging_face_nmt_model_trainer import HuggingFaceNmtModelTrainer, add_lang_code_to_tokenizer

__all__ = ["add_lang_code_to_tokenizer", "HuggingFaceNmtEngine", "HuggingFaceNmtModel", "HuggingFaceNmtModelTrainer"]
__all__ = [
"add_lang_code_to_tokenizer",
"HuggingFaceNmtEngine",
"HuggingFaceNmtModel",
"HuggingFaceNmtModelTrainer",
"SilTranslationPipeline",
]
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
{
"additional_special_tokens": null,
"extra_special_tokens": null,
"bos_token": "<s>",
"cls_token": "<s>",
"eos_token": "</s>",
Expand Down
Loading
Loading