Skip to content
Open
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
56 changes: 42 additions & 14 deletions train_chandra.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
# ---------------------------------------------------------------------------
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -741,13 +759,18 @@ 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)
log.info("PHASE: Pre-training inference test")
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")
Expand Down Expand Up @@ -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:
Expand Down