From b7002da391339d246f9d6f00c3bd3edb563ad63c Mon Sep 17 00:00:00 2001 From: Vedant Kambli Date: Sun, 31 May 2026 15:07:36 +0530 Subject: [PATCH] Fix dataset map multiprocessing and PIL image byte loading --- train_chandra.py | 56 ++++++++++++++++++++++++++++++++++++------------ 1 file changed, 42 insertions(+), 14 deletions(-) diff --git a/train_chandra.py b/train_chandra.py index 5bfdb89..6e255d8 100644 --- a/train_chandra.py +++ b/train_chandra.py @@ -32,6 +32,11 @@ from pathlib import Path from typing import Any, Dict, List, Optional, Union + + +import io +from PIL import Image + # --------------------------------------------------------------------------- # 2. Configuration Constants # --------------------------------------------------------------------------- @@ -444,26 +449,39 @@ def prepare_dataset( ) fmt = "simple" - def _convert_split(dataset): - return [ - convert_to_conversation( - s, - instruction=args.ocr_instruction, - image_col=args.image_column, - text_col=args.text_column, - dataset_format=fmt, - ) - for s in dataset - ] + def _map_sample(sample): + return convert_to_conversation( + sample, + instruction=args.ocr_instruction, + image_col=args.image_column, + text_col=args.text_column, + dataset_format=fmt, + ) log.info("Converting %d train samples to conversation format...", len(train_ds)) - train_conv = _convert_split(train_ds) + + workers = min(8, os.cpu_count() or 1) + + print(f"Using {workers} cores to load the datasets... ") + train_conv = train_ds.map( + _map_sample, + remove_columns=train_ds.column_names, + desc="Formatting train set", + num_proc = workers, + writer_batch_size=75 + ) eval_conv = None if eval_ds is not None: log.info("Converting %d eval samples to conversation format...", len(eval_ds)) - eval_conv = _convert_split(eval_ds) - + eval_conv = eval_ds.map( + _map_sample, + remove_columns=eval_ds.column_names, + desc="Formatting eval set", + num_proc = workers, + writer_batch_size=75 + ) + log.info("Conversion complete: train=%d eval=%s", len(train_conv), len(eval_conv) if eval_conv else "None") return train_conv, eval_conv @@ -741,6 +759,8 @@ def train(args: argparse.Namespace) -> None: log.info("=" * 60) train_data, eval_data = prepare_dataset(args) + print("==================== Dataset prepared ====================") + # -- 5. Pre-training inference test -- if not args.skip_pre_eval: log.info("=" * 60) @@ -748,6 +768,9 @@ def train(args: argparse.Namespace) -> None: log.info("=" * 60) _run_inference_test(model, tokenizer, processor, train_data, tag="BEFORE training") + + print("==================== pre-trainig inference done ====================") + # -- 6. Train -- log.info("=" * 60) log.info("PHASE: Training") @@ -828,6 +851,11 @@ def _run_inference_test( for part in user_content: if part.get("type") == "image": test_image = part.get("image") + + # Handle datasets where images are returned as raw bytes dictionaries + if isinstance(test_image, dict) and "bytes" in test_image: + test_image = Image.open(io.BytesIO(test_image["bytes"])).convert("RGB") + break if test_image is None: