ml-explore / ml-explore/mlx-examples

iterate_batches in mlx_lm's Lora trainer is discarding the remainder dataset items (modulo batch size)

Open
#843 1 comment 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

enhancement
Dominant language
Python
Stars
9k
Forks
1.2k
PR merge metrics
No merged PRs in 30d

Description

The current implementation of iterate_batches produces batches for all but the remaining N items in the dataset (where N is less than the batch size). So, if your dataset size is a multiple of the batch size, you will eventually train on every item in the dataset. However, if it is not, the remainder will never be included in what is trained, regardless of how many iterations you set.

Below is very minimal test case that replicates this (it generates batches of a total size of at most 46 items without finding the remaining data before stopping):

from unittest.mock import MagicMock
from mlx_lm.tuner.trainer import iterate_batches
import mlx.core as mx

class TestBatching(unittest.TestCase):
    def test_batch_remainder(self):
        tokenizer = MagicMock()
        tokenizer.eos_token_id = 9
        tokenizer.encode = lambda x: [1,] if x == "foo" else [7,]
        
        #The dataset is 20 "foo"s and 3 "bar"s (the remainder) 
        dataset = ["foo"] * 20 + ["bar", "bar", "bar"]
        batch_size = 10
        found_remainders = False
        for idx, (batch_in, _, _) in enumerate(iterate_batches(dataset, tokenizer, batch_size, max_seq_length=2048,
                                                               train=True)):
            found_remainders = mx.any(batch_in == mx.full(batch_in.shape, 7))
            if found_remainders or ((idx + 1) * batch_size) > len(dataset) * 2:
                break
        self.assertTrue(found_remainders.item())

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start with iterate_batches in mlx_lm/tuner/trainer.py and reproduce the issue using the provided TestBatching case with a 23-item dataset and batch size 10. Done means the generated batches eventually include the three remainder items, and the regression test passes.

Written by the indexing model from the issue text.

Assessment

Tech stack
python
Domain
machine-learning
Issue type
Bug
Difficulty
2/5
Estimated time
1-3 hours
Activity status
Stale
Clarity
Clearly specified
Newbie friendliness
55/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.