olmOCR-mix-1025 datat may ber something wrong
- Lingua principale
- Python
- Stelle
- 19.5k
- Fork
- 1.6k
- Metriche di merge delle PR
- Nessuna PR unita negli ultimi 30g
Descrizione
### 🐛 Describe the bug
download olmOCR-mix-1025 and run prepare_olmocrmix to extract data,then i just use processed_00_documents_eval_s2pdf for test training
"""
Simple script to test OlmOCR dataset loading with YAML configuration.
"""
import argparse
import logging
import os
from typing import Optional
from dataclasses import dataclass
from transformers import DataCollatorWithPadding,DataCollatorForSeq2Seq
from typing import Any
import numpy as np
import torch
from torch.utils.data import ConcatDataset
from transformers import (
AutoProcessor,
EarlyStoppingCallback,
Qwen2_5_VLForConditionalGeneration,
Qwen2VLForConditionalGeneration,
Trainer,
TrainingArguments,
)
from transformers import TrainingArguments
from olmocr.train.config import Config
from olmocr.train.dataloader import BaseMarkdownPDFDataset
# Configure logging
logging.basicConfig(
format="%(asctime)s - %(levelname)s - %(name)s - %(message)s",
datefmt="%m/%d/%Y %H:%M:%S",
level=logging.INFO,
)
logger = logging.getLogger(__name__)
@dataclass
class QwenvlDatacollator(DataCollatorForSeq2Seq):
def __call__(self, features):
try:
feature_test = [
{k: v for k, v in feature.items() if k in ["input_ids", "attention_mask", "labels"]}
for feature in features
]
except:
return None
#self.verbose_list.extend(features)
# Use the parent class to handle padding for input_ids, attention_mask, and labels
batch = super().__call__(feature_test)
#for pixes in features:
for feature in features:
if "pixel_values" in feature:
if "pixel_values" not in batch:
batch["pixel_values"] = []
batch["pixel_values"].append(feature["pixel_values"])
if "image_grid_thw" in feature:
if "image_grid_thw" not in batch:
batch["image_grid_thw"] = []
batch["image_grid_thw"].append(feature["image_grid_thw"])
return {
"input_ids":batch["input_ids"],
"attention_mask": batch["attention_mask"],
"labels": batch["labels"],
"pixel_values": torch.vstack(batch["pixel_values"]), # Stack into tensor
"image_grid_thw": torch.stack(batch["image_grid_thw"]),
}
def main():
parser = argparse.ArgumentParser(description="Train OlmOCR model")
parser.add_argument("--config", type=str, default="olmocr/train/configs/v0.4.0/qwen25_vl_olmocrv4_rotation_1epoch_mix_1025_filtered.yaml", help="Path to YAML configuration file")
args = parser.parse_args()
# Load configuration
logger.info(f"Loading configuration from: {args.config}")
config = Config.from_yaml(args.config)
# Validate configuration
try:
config.validate()
except ValueError as e:
logger.error(f"Configuration validation failed: {e}")
return
# Set wandb project from config
if config.project_name:
os.environ["WANDB_PROJECT"] = config.project_name
logger.info(f"Setting WANDB_PROJECT to: {config.project_name}")
# Load processor for tokenization
logger.info(f"Loading processor: {config.model.name}")
processor = AutoProcessor.from_pretrained(
config.model.name,
)
# Load model
logger.info(f"Loading model: {config.model.name}")
if "Qwen2.5-VL" in config.model.name:
model = Qwen2_5_VLForConditionalGeneration.from_pretrained(
config.model.name,
torch_dtype=getattr(torch, config.model.torch_dtype) if config.model.torch_dtype != "auto" else "auto",
#device_map=config.model.device_map,
trust_remote_code=config.model.trust_remote_code,
attn_implementation=config.model.attn_implementation if config.model.use_flash_attention else None,
)
elif "Qwen2-VL" in config.model.name:
model = Qwen2VLForConditionalGeneration.from_pretrained(
config.model.name,
torch_dtype=getattr(torch, config.model.torch_dtype) if config.model.torch_dtype != "auto" else "auto",
#device_map=config.model.device_map,
trust_remote_code=config.model.trust_remote_code,
attn_implementation=config.model.attn_implementation if config.model.use_flash_attention else None,
)
else:
raise NotImplementedError()
# Enable gradient checkpointing if configured
if config.training.gradient_checkpointing:
model.gradient_checkpointing_enable(gradient_checkpointing_kwargs=config.training.gradient_checkpointing_kwargs)
# Create training datasets
logger.info("Creating training datasets...")
train_datasets = []
for i, dataset_cfg in enumerate(config.dataset.train):
root_dir = dataset_cfg["root_dir"]
pipeline_steps = config.get_pipeline_steps(dataset_cfg["pipeline"], processor)
logger.info(f"Creating training dataset {i+1} from: {root_dir}")
dataset = BaseMarkdownPDFDataset(root_dir, pipeline_steps)
logger.info(f"Found {len(dataset)} samples")
if len(dataset) > 0:
train_datasets.append(dataset)
# Combine all training datasets
train_dataset = ConcatDataset(train_datasets) if len(train_datasets) > 1 else train_datasets[0]
logger.info(f"Total training samples: {len(train_dataset)}")
# Create evaluation datasets
logger.info("Creating evaluation datasets...")
eval_datasets = {}
for i, dataset_cfg in enumerate(config.dataset.eval):
root_dir = dataset_cfg["root_dir"]
pipeline_steps = config.get_pipeline_steps(dataset_cfg["pipeline"], processor)
# Use dataset name if provided, otherwise use root_dir as name
dataset_name = dataset_cfg.get("name", f"eval_dataset_{i+1}")
logger.info(f"Creating evaluation dataset '{dataset_name}' from: {root_dir}")
dataset = BaseMarkdownPDFDataset(root_dir, pipeline_steps)
logger.info(f"Found {len(dataset)} samples")
if len(dataset) > 0:
eval_datasets[dataset_name] = dataset
# Log total evaluation samples across all datasets
total_eval_samples = sum(len(dataset) for dataset in eval_datasets.values())
logger.info(f"Total evaluation samples across {len(eval_datasets)} datasets: {total_eval_samples}")
# Construct full output directory by appending run_name to base output_dir
full_output_dir = os.path.join(config.training.output_dir, config.run_name)
logger.info(f"Setting output directory to: {full_output_dir}")
# Check for existing checkpoints if any
found_resumable_checkpoint = False
if os.path.exists(full_output_dir):
# Look for checkpoint directories
checkpoint_dirs = [d for d in os.listdir(full_output_dir) if d.startswith("checkpoint-") and os.path.isdir(os.path.join(full_output_dir, d))]
if checkpoint_dirs:
# Sort by checkpoint number and get the latest
checkpoint_dirs.sort(key=lambda x: int(x.split("-")[1]))
latest_checkpoint = os.path.join(full_output_dir, checkpoint_dirs[-1])
logger.info(f"Found existing checkpoint: {latest_checkpoint}")
found_resumable_checkpoint = latest_checkpoint
else:
logger.info("No existing checkpoints found in output directory")
# Set up training arguments
training_args = TrainingArguments(
output_dir=full_output_dir,
num_train_epochs=config.training.num_train_epochs,
per_device_train_batch_size=config.training.per_device_train_batch_size,
per_device_eval_batch_size=config.training.per_device_eval_batch_size,
gradient_accumulation_steps=config.training.gradient_accumulation_steps,
learning_rate=float(config.training.learning_rate),
lr_scheduler_type=config.training.lr_scheduler_type,
warmup_ratio=config.training.warmup_ratio,
lr_scheduler_kwargs=config.training.lr_scheduler_kwargs,
optim=config.training.optim,
adam_beta1=config.training.adam_beta1,
adam_beta2=config.training.adam_beta2,
adam_epsilon=config.training.adam_epsilon,
weight_decay=config.training.weight_decay,
max_grad_norm=config.training.max_grad_norm,
bf16=True, # We're sticking with this known good reduced precision option
eval_strategy=config.training.evaluation_strategy,
eval_steps=config.training.eval_steps,
save_strategy=config.training.save_strategy,
save_steps=config.training.save_steps,
save_total_limit=config.training.save_total_limit,
load_best_model_at_end=config.training.load_best_model_at_end,
metric_for_best_model=config.training.metric_for_best_model,
greater_is_better=config.training.greater_is_better,
logging_dir=config.training.logging_dir,
logging_strategy=config.training.logging_strategy,
logging_steps=config.training.logging_steps,
logging_first_step=config.training.logging_first_step,
report_to=config.training.report_to,
seed=config.training.seed,
data_seed=config.training.data_seed,
push_to_hub=False,
label_names=["labels"],
dataloader_drop_last=config.training.dataloader_drop_last,
dataloader_num_workers=config.training.dataloader_num_workers,
remove_unused_columns=config.training.remove_unused_columns,
eval_on_start=True,
run_name=config.run_name,
)
# Set up callbacks
callbacks = []
if config.training.use_early_stopping:
callbacks.append(
EarlyStoppingCallback(
early_stopping_patience=config.training.early_stopping_patience, early_stopping_threshold=config.training.early_stopping_threshold
)
)
data_collator = QwenvlDatacollator(max_length=config.training.collator_max_token_len, tokenizer=processor.tokenizer, model=model,padding="max_length")
# Initialize trainer
logger.info("Initializing trainer...")
trainer = Trainer(
model=model,
args=training_args,
train_dataset=train_dataset,
eval_dataset=eval_datasets,
data_collator=data_collator,
callbacks=callbacks,
)
# Start training
logger.info("Starting training...")
train_result = trainer.train(resume_from_checkpoint=found_resumable_checkpoint)
# Save the final model
logger.info("Saving final model...")
trainer.save_model()
trainer.save_state()
# Log metrics
logger.info(f"Training completed! Metrics: {train_result.metrics}")
if __name__ == "__main__":
main()
config:
# Example OlmOCR Training Configuration with Torch Compile
# Project metadata
project_name: olmocr-qwen-vl-training
run_name: qwen2.5-vl-7b-olmocrv4_1epoch_promptv4_mix1025_more_rotation_filtered
# Model configuration
model:
name: QWEN2/Qwen2.5-VL-3B-Instruct
trust_remote_code: true
torch_dtype: bfloat16
use_flash_attention: true
attn_implementation: flash_attention_2
# LoRA settings (disabled by default)
use_lora: false
# lora_rank: 8
# lora_alpha: 32
# lora_dropout: 0.1
# lora_target_modules:
# - q_proj
# - v_proj
# - k_proj
# - o_proj
# Dataset configuration
dataset:
train:
- name: processed_00_documents_eval_s2pdf
root_dir: olmOCR-mix-1025-extracted/processed_00_documents_eval/
pipeline: &basic_pipeline
- name: FrontMatterParser
front_matter_class: PageResponse
- name: FilterOutRotatedDocuments
- name: ReformatLatexBoldItalic
- name: DatasetTextRuleFilter
- name: PDFRenderer
target_longest_image_dim: 1288
- name: RotationAugmentation
probability: 0.02
- name: NewYamlFinetuningPromptWithNoAnchoring
- name: FrontMatterOutputFormat
- name: InstructUserMessages
prompt_first: true
- name: Tokenizer
masking_index: -100
end_of_message_token: "<|im_end|>"
# - name: processed_00_documents_train
# root_dir: olmOCR-mix-1025-extracted/processed_00_documents_train/
# pipeline: *basic_pipeline
# - name: processed_02_loc_transcripts_train
# root_dir: olmOCR-mix-1025-extracted/processed_02_loc_transcripts_train/
# pipeline: *basic_pipeline
# - name: processed_03_national_archives
# root_dir: olmOCR-mix-1025-extracted/processed_03_national_archives_train/
# pipeline: *basic_pipeline
eval:
- name: processed_00_documents_eval_s2pdf
root_dir: olmOCR-mix-1025-extracted/processed_00_documents_eval/
pipeline: *basic_pipeline
# - name: processed_01_books_eval_iabooks
# root_dir: olmOCR-mix-1025-extracted/processed_01_books_eval/
# pipeline: *basic_pipeline
# - name: processed_02_loc_transcripts_eval
# root_dir: olmOCR-mix-1025-extracted/processed_02_loc_transcripts_eval/
# pipeline: *basic_pipeline
# - name: processed_03_national_archives_eval
# root_dir: olmOCR-mix-1025-extracted/processed_03_national_archives_eval/
# pipeline: *basic_pipeline
# Training configuration
training:
output_dir: ./olmocr-trainer/
num_train_epochs: 1
# Batch size and accumulation
per_device_train_batch_size: 1
per_device_eval_batch_size: 1
gradient_accumulation_steps: 32
gradient_checkpointing: False
collator_max_token_len: 8192
# Learning rate
learning_rate: 2e-5
lr_scheduler_type: linear
warmup_ratio: 0.1
# Optimization
optim: adamw_torch
weight_decay: 0.01
max_grad_norm: 1.0
# Torch compile settings
torch_compile: false
torch_compile_backend: inductor
torch_compile_mode: default
torch_compile_fullgraph: false
torch_compile_dynamic: false
seed: 300
data_seed: 301
# dataloader_num_workers: 8
# Evaluation and checkpointing
evaluation_strategy: steps
eval_steps: 500
save_strategy: steps
save_steps: 500
save_total_limit: 2
load_best_model_at_end: false # Needs to be false because it has a problem restoring checkpoints for some reason
metric_for_best_model: eval_processed_00_documents_eval_s2pdf_loss
greater_is_better: false
report_to:
- wandb
error:
return inner_training_loop(
^^^^^^^^^^^^^^^^^^^^
File "/lib/python3.12/site-packages/transformers/trainer.py", line 2576, in _inner_training_loop
self._evaluate(trial, ignore_keys_for_eval, skip_scheduler=True)
File "/lib/python3.12/site-packages/transformers/trainer.py", line 3170, in _evaluate
metrics = self.evaluate(ignore_keys=ignore_keys_for_eval)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/lib/python3.12/site-packages/transformers/trainer.py", line 4471, in evaluate
dataset_metrics = self.evaluate(
^^^^^^^^^^^^^^
File "/lib/python3.12/site-packages/transformers/trainer.py", line 4489, in evaluate
output = eval_loop(
^^^^^^^^^^
File "/lib/python3.12/site-packages/transformers/trainer.py", line 4685, in evaluation_loop
losses, logits, labels = self.prediction_step(model, inputs, prediction_loss_only, ignore_keys=ignore_keys)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/lib/python3.12/site-packages/transformers/trainer.py", line 4854, in prediction_step
has_labels = False if len(self.label_names) == 0 else all(inputs.get(k) is not None for k in self.label_names)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/lib/python3.12/site-packages/transformers/trainer.py", line 4854, in
has_labels = False if len(self.label_names) == 0 else all(inputs.get(k) is not None for k in self.label_names)
^^^^^^^^^^
AttributeError: 'NoneType' object has no attribute 'get'
[rank0]: Traceback (most recent call last):
[rank0]: File "/home/user/chatglm/llm_ki/olmocr/olmocr/train/train_mult_gpu.py", line 313, in
[rank0]: main()
[rank0]: File "/home/user/chatglm/llm_ki/olmocr/olmocr/train/train_mult_gpu.py", line 301, in main
[rank0]: train_result = trainer.train(resume_from_checkpoint=found_resumable_checkpoint)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/lib/python3.12/site-packages/transformers/trainer.py", line 2325, in train
[rank0]: return inner_training_loop(
[rank0]: ^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/lib/python3.12/site-packages/transformers/trainer.py", line 2576, in _inner_training_loop
[rank0]: self._evaluate(trial, ignore_keys_for_eval, skip_scheduler=True)
[rank0]: File "/lib/python3.12/site-packages/transformers/trainer.py", line 3170, in _evaluate
[rank0]: metrics = self.evaluate(ignore_keys=ignore_keys_for_eval)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/lib/python3.12/site-packages/transformers/trainer.py", line 4471, in evaluate
[rank0]: dataset_metrics = self.evaluate(
[rank0]: ^^^^^^^^^^^^^^
[rank0]: File "/lib/python3.12/site-packages/transformers/trainer.py", line 4489, in evaluate
[rank0]: output = eval_loop(
[rank0]: ^^^^^^^^^^
[rank0]: File "/lib/python3.12/site-packages/transformers/trainer.py", line 4685, in evaluation_loop
[rank0]: losses, logits, labels = self.prediction_step(model, inputs, prediction_loss_only, ignore_keys=ignore_keys)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/lib/python3.12/site-packages/transformers/trainer.py", line 4854, in prediction_step
[rank0]: has_labels = False if len(self.label_names) == 0 else all(inputs.get(k) is not None for k in self.label_names)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/lib/python3.12/site-packages/transformers/trainer.py", line 4854, in
[rank0]: has_labels = False if len(self.label_names) == 0 else all(inputs.get(k) is not None for k in self.label_names)
[rank0]: ^^^^^^^^^^
[rank0]: AttributeError: 'NoneType' object has no attribute 'get'
wandb:
wandb: 🚀 View run qwen2.5-vl-7b-olmocrv4_1epoch_promptv4_mix1025_more_rotation_filtered at:
### Versions
olmocr 0.4.2
Guida per i contributori
Apri la guida per i contributori
Valutazione
Questa issue non è ancora stata valutata.