MultiModal / data_loader.py
szxllm's picture
Update data_loader.py
b7e4db5 verified
Raw
History Blame Contribute Delete
27.8 kB
import torch
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader, IterableDataset
from datasets import load_dataset, concatenate_datasets, interleave_datasets
from typing import Dict, List, Optional, Any, Union
import random
import numpy as np
from tqdm import tqdm
import warnings
from PIL import Image
import requests
from io import BytesIO
from torchvision import transforms
import logging
import os
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
warnings.filterwarnings("ignore", category=UserWarning)
from data_config import (
PRETRAIN_DATASETS,
POSTTRAIN_DATASETS,
TEST_DATASETS,
PRETRAIN_MIX,
POSTTRAIN_MIX,
PREPROCESSING_CONFIG,
DATASET_CACHE_DIR,
HF_CACHE_DIR
)
image_transform = transforms.Compose([
transforms.Resize((224, 224)),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])
class PreTrainDataset(IterableDataset):
def __init__(
self,
mix_name: str = 'default',
tokenizer=None,
max_length: int = 2048,
streaming: bool = True,
seed: int = 42,
max_samples: Optional[int] = None
):
super().__init__()
if tokenizer is None:
raise ValueError("tokenizer cannot be None")
self.tokenizer = tokenizer
self.max_length = max_length
self.streaming = streaming
self.seed = seed
self.max_samples = max_samples
self.samples_generated = 0
if mix_name not in PRETRAIN_MIX:
raise ValueError(f"Unknown mix: {mix_name}. Available: {list(PRETRAIN_MIX.keys())}")
mix_config = PRETRAIN_MIX[mix_name]
dataset_names = mix_config.get('datasets', [])
weights = mix_config.get('weights', [])
if not dataset_names:
raise ValueError(f"No datasets found in mix: {mix_name}")
self.datasets = []
self.probabilities = []
for name, weight in zip(dataset_names, weights):
if name not in PRETRAIN_DATASETS:
logger.warning(f"Dataset {name} not found in PRETRAIN_DATASETS, skipping")
continue
config = PRETRAIN_DATASETS[name]
try:
ds = self._load_dataset(config)
if ds is not None:
self.datasets.append((name, ds, config))
self.probabilities.append(weight)
logger.info(f" Successfully loaded {name}")
except Exception as e:
logger.error(f"Error loading {name}: {e}")
continue
if not self.datasets:
raise ValueError("No datasets loaded successfully")
total = sum(self.probabilities)
self.probabilities = [p / total for p in self.probabilities]
logger.info(f"Successfully loaded {len(self.datasets)} datasets")
def _load_dataset(self, config: Dict):
try:
load_kwargs = {
'path': config['hf_path'],
'split': config.get('split', 'train'),
'streaming': config.get('streaming', self.streaming),
'cache_dir': HF_CACHE_DIR,
}
if 'data_files' in config:
files = config['data_files']
if isinstance(files, list):
for f in files:
if not os.path.exists(f):
logger.error(f" Data file not found in list: {f}")
return None
logger.info(f" Verified {len(files)} local files.")
elif isinstance(files, str):
if not os.path.exists(files):
logger.error(f" Data file not found: {files}")
return None
logger.info(f" Verified local file: {files}")
load_kwargs['data_files'] = files
if 'config' in config:
load_kwargs['name'] = config['config']
logger.info(f" Loading HF dataset: {config['hf_path']}...")
ds = load_dataset(**load_kwargs)
return ds
except Exception as e:
logger.error(f"Failed to load {config.get('hf_path', 'unknown')}: {e}")
import traceback
traceback.print_exc()
return None
def _process_text_sample(self, sample: Dict, config: Dict) -> Optional[Dict]:
try:
text_field = config.get('text_field', 'text')
text = sample.get(text_field, '')
if not text or not isinstance(text, str):
return None
text = text.strip()
if len(text) < 10:
return None
max_input_len = self.max_length - 1
encoding = self.tokenizer(
text,
max_length=max_input_len,
truncation=True,
padding=False,
add_special_tokens=False,
return_tensors=None
)
input_ids = encoding['input_ids']
input_ids.append(self.tokenizer.eos_token_id)
input_ids_tensor = torch.tensor(input_ids, dtype=torch.long)
return {
'input_ids': input_ids_tensor,
'type': 'text'
}
except Exception as e:
logger.debug(f"Error processing text sample: {e}")
return None
def _process_image_text_sample(self, sample: Dict, config: Dict) -> Optional[Dict]:
try:
text_field = config.get('text_field', 'caption')
image_field = config.get('image_field', 'image')
text = sample.get(text_field, '')
image = sample.get(image_field)
if not text or image is None:
return None
if isinstance(image, str):
try:
response = requests.get(image, timeout=5)
image = Image.open(BytesIO(response.content)).convert('RGB')
except Exception as img_error:
logger.debug(f"Failed to load image from URL: {img_error}")
return None
elif isinstance(image, Image.Image):
image = image.convert('RGB')
else:
return None
image_tensor = image_transform(image)
encoding = self.tokenizer(
text,
max_length=self.max_length,
truncation=True,
padding='max_length',
return_tensors='pt'
)
return {
'input_ids': encoding['input_ids'].squeeze(0),
'attention_mask': encoding['attention_mask'].squeeze(0),
'image': image_tensor,
'type': 'image_text'
}
except Exception as e:
logger.debug(f"Error processing image-text sample: {e}")
return None
def __iter__(self):
worker_info = torch.utils.data.get_worker_info()
if worker_info is not None:
random.seed(self.seed + worker_info.id)
np.random.seed(self.seed + worker_info.id)
else:
random.seed(self.seed)
np.random.seed(self.seed)
iterators = [iter(ds) for _, ds, _ in self.datasets]
self.samples_generated = 0
while True:
if self.max_samples and self.samples_generated >= self.max_samples:
break
try:
idx = np.random.choice(len(self.datasets), p=self.probabilities)
name, _, config = self.datasets[idx]
sample = next(iterators[idx])
processed = None
if config.get('type') in ['text', 'code']:
processed = self._process_text_sample(sample, config)
elif config.get('type') == 'image_text':
processed = self._process_image_text_sample(sample, config)
else:
logger.debug(f"Unknown type: {config.get('type')}")
continue
if processed is not None:
self.samples_generated += 1
yield processed
except StopIteration:
try:
iterators[idx] = iter(self.datasets[idx][1])
except Exception as e:
logger.error(f"Failed to recreate iterator for dataset {idx}: {e}")
break
except Exception as e:
logger.debug(f"Error in iterator: {e}")
continue
class PostTrainDataset(Dataset):
def __init__(
self,
mix_name: str = 'default',
tokenizer=None,
max_length: int = 2048,
max_samples: Optional[int] = None,
split: str = 'train'
):
super().__init__()
if tokenizer is None:
raise ValueError("tokenizer cannot be None")
self.tokenizer = tokenizer
self.max_length = max_length
self.split = split
if mix_name not in POSTTRAIN_MIX:
raise ValueError(f"Unknown mix: {mix_name}. Available: {list(POSTTRAIN_MIX.keys())}")
mix_config = POSTTRAIN_MIX[mix_name]
dataset_names = mix_config.get('datasets', [])
weights = mix_config.get('weights', [])
if not dataset_names:
raise ValueError(f"No datasets found in mix: {mix_name}")
logger.info(f"Loading posttrain mix: {mix_name}")
logger.info(f" Datasets: {dataset_names}")
all_datasets = []
for name in dataset_names:
if name not in POSTTRAIN_DATASETS:
logger.warning(f"Dataset {name} not found in POSTTRAIN_DATASETS")
continue
config = POSTTRAIN_DATASETS[name]
try:
load_kwargs = {
'path': config['hf_path'],
'split': split,
'streaming': config.get('streaming', False),
'cache_dir': HF_CACHE_DIR,
}
if 'data_files' in config:
load_kwargs['data_files'] = config['data_files']
if 'config' in config:
load_kwargs['name'] = config['config']
ds = load_dataset(**load_kwargs)
if config.get('max_samples'):
if hasattr(ds, 'take'):
ds = ds.take(config['max_samples'])
elif hasattr(ds, 'select'):
ds = ds.select(range(min(len(ds), config['max_samples'])))
def add_source(example):
example['_source'] = name
example['_config'] = config
return example
ds = ds.map(add_source)
all_datasets.append(ds)
ds_len = len(ds) if hasattr(ds, '__len__') else 'streaming'
logger.info(f" Loaded {name}: {ds_len} samples")
except Exception as e:
logger.error(f"Error loading {name}: {e}")
continue
if not all_datasets:
raise ValueError("No datasets loaded successfully")
if len(all_datasets) == 1:
self.dataset = all_datasets[0]
else:
probabilities = [w / sum(weights[:len(all_datasets)])
for w in weights[:len(all_datasets)]]
self.dataset = interleave_datasets(
all_datasets,
probabilities=probabilities,
seed=42,
stopping_strategy='all_exhausted'
)
if max_samples and hasattr(self.dataset, '__len__'):
actual_len = min(len(self.dataset), max_samples)
self.dataset = self.dataset.select(range(actual_len))
dataset_len = len(self.dataset) if hasattr(self.dataset, '__len__') else 'streaming'
logger.info(f"Total samples: {dataset_len}")
def _format_instruction(self, sample: Dict, config: Dict) -> str:
try:
data_type = config.get('type', 'instruction')
if data_type == 'instruction':
instruction_field = config.get('instruction_field', 'instruction')
input_field = config.get('input_field', 'input')
context_field = config.get('context_field', None)
instruction = sample.get(instruction_field, '')
input_text = sample.get(input_field, '')
context = sample.get(context_field, '') if context_field else ''
prompt_parts = [f"Instruction: {instruction}"]
if context:
prompt_parts.append(f"Context: {context}")
if input_text:
prompt_parts.append(f"Input: {input_text}")
prompt_parts.append("Response:")
return "\n".join(prompt_parts)
elif data_type == 'conversation':
if 'conversations' in sample:
conversations = sample['conversations']
if isinstance(conversations, list) and len(conversations) > 0:
dialogue = []
last_role = conversations[-1].get('role', conversations[-1].get('from', 'user')).lower()
upto = len(conversations)
if last_role == 'assistant':
upto = len(conversations) - 1
for conv in conversations[:upto]:
role = conv.get('role', conv.get('from', 'user'))
content = conv.get('content', conv.get('value', ''))
dialogue.append(f"{role}: {content}")
return "\n".join(dialogue) + "\nassistant:"
elif 'messages' in sample:
messages = sample['messages']
if isinstance(messages, list) and len(messages) > 0:
dialogue = []
last_role = messages[-1].get('role', 'user').lower()
upto = len(messages)
if last_role == 'assistant':
upto = len(messages) - 1
for msg in messages[:upto]:
role = msg.get('role', 'user')
content = msg.get('content', '')
dialogue.append(f"{role}: {content}")
return "\n".join(dialogue) + "\nassistant:"
return sample.get('text', '')
elif data_type == 'code_instruction':
instruction_field = config.get('instruction_field', 'instruction')
instruction = sample.get(instruction_field, '')
return f"### Instruction:\n{instruction}\n### Response:"
elif data_type == 'multimodal_instruction':
instruction_field = config.get('instruction_field', 'conversations')
conversations = sample.get(instruction_field, [])
if isinstance(conversations, list) and len(conversations) > 0:
dialogue = []
last_role = conversations[-1].get('from', 'user').lower() if isinstance(conversations[-1].get('from', 'user'), str) else 'user'
upto = len(conversations)
if last_role == 'assistant':
upto = len(conversations) - 1
for conv in conversations[:upto]:
role = conv.get('from', 'user')
content = conv.get('value', '')
dialogue.append(f"{role}: {content}")
return "\n".join(dialogue) + "\nassistant:"
return ""
else:
return sample.get(config.get('instruction_field', 'text'), '')
except Exception as e:
logger.debug(f"Error formatting instruction: {e}")
return ""
def _get_response(self, sample: Dict, config: Dict) -> str:
"""获取响应(兼容 <think>/<answer> 标签)"""
try:
data_type = config.get('type', 'instruction')
if data_type == 'instruction' or data_type == 'code_instruction':
response_field = config.get('response_field', 'output')
return sample.get(response_field, '')
elif data_type == 'conversation':
# 从对话中提取最后一条assistant的回复
if 'conversations' in sample:
conversations = sample['conversations']
if isinstance(conversations, list) and len(conversations) > 0:
last_turn = conversations[-1]
content = last_turn.get('content', last_turn.get('value', ''))
if not isinstance(content, str):
return ''
# 仅当最后一条 role 为 assistant 时返回
role = last_turn.get('role', last_turn.get('from', '')).lower()
if role != 'assistant':
return ''
return str(content).strip() if content else ""
elif 'messages' in sample:
messages = sample['messages']
if isinstance(messages, list) and len(messages) > 0:
return messages[-1].get('content', '')
return ""
elif data_type == 'multimodal_instruction':
instruction_field = config.get('instruction_field', 'conversations')
conversations = sample.get(instruction_field, [])
if isinstance(conversations, list) and len(conversations) > 0:
last = conversations[-1].get('value', '')
import re
m = re.search(r'<answer>([\\s\\S]*?)</answer>', last, re.IGNORECASE)
if m:
return m.group(1).strip()
return re.sub(r'<think>[\\s\\S]*?</think>', '', last, flags=re.IGNORECASE).strip()
return ""
else:
response_field = config.get('response_field', 'output')
return sample.get(response_field, '')
except Exception as e:
logger.debug(f"Error getting response: {e}")
return ""
def __len__(self):
return len(self.dataset) if hasattr(self.dataset, '__len__') else 0
def __getitem__(self, idx):
try:
sample = self.dataset[idx]
if '_config' not in sample:
logger.warning(f"Sample at index {idx} missing _config")
return None
config = sample['_config']
instruction_text = self._format_instruction(sample, config)
response_text = self._get_response(sample, config)
if not instruction_text or not response_text:
return None
pad_token_id = self.tokenizer.pad_token_id
if pad_token_id is None:
pad_token_id = self.tokenizer.eos_token_id
instruction_max_len = 256
instruction_enc = self.tokenizer(
instruction_text,
truncation=True,
max_length=instruction_max_len,
add_special_tokens=False,
return_tensors=None
)
instr_ids_list = instruction_enc['input_ids']
instr_ids = torch.tensor(instr_ids_list, dtype=torch.long)
response_max_len = self.max_length - len(instr_ids)
response_enc = self.tokenizer(
response_text,
truncation=True,
max_length=response_max_len - 1,
add_special_tokens=False,
return_tensors=None
)
resp_ids_list = response_enc['input_ids']
resp_ids_list = resp_ids_list + [self.tokenizer.eos_token_id]
resp_ids = torch.tensor(resp_ids_list, dtype=torch.long)
result = {
'instruction': instr_ids,
'response': resp_ids,
'task': sample.get('_source', 'unknown'),
'modality_data': None
}
if config.get('type') == 'multimodal_instruction' and 'image' in sample:
try:
image = sample['image']
if isinstance(image, Image.Image):
image = image.convert('RGB')
image_tensor = image_transform(image)
result['modality_data'] = {'image': image_tensor}
except Exception as e:
logger.debug(f"Error processing image: {e}")
return result
except Exception as e:
logger.debug(f"Error getting item at index {idx}: {e}")
import traceback
traceback.print_exc()
return None
from torch.nn.utils.rnn import pad_sequence
class DynamicCollate:
def __init__(self, pad_token_id: int):
self.pad_token_id = pad_token_id
def __call__(self, batch):
batch = [item for item in batch if item is not None]
if not batch:
return {
'input_ids': torch.empty(0),
'attention_mask': torch.empty(0)
}
input_ids_list = [item['input_ids'] for item in batch]
padded_input_ids = pad_sequence(
input_ids_list,
batch_first=True,
padding_value=self.pad_token_id
)
attention_mask = (padded_input_ids != self.pad_token_id).long()
return {
'input_ids': padded_input_ids,
'attention_mask': attention_mask
}
def collate_fn_v2_factory(pad_token_id: int):
def collate_fn_v2(batch):
batch = [item for item in batch if item is not None]
if not batch:
logger.warning("Empty batch after filtering None values")
return None
if isinstance(batch[0], tuple):
if len(batch[0]) == 4:
chosen = torch.stack([item[0] for item in batch])
rejected = torch.stack([item[1] for item in batch])
chosen_mask = torch.stack([item[2] for item in batch])
rejected_mask = torch.stack([item[3] for item in batch])
return {
'chosen': chosen,
'rejected': rejected,
'chosen_mask': chosen_mask,
'rejected_mask': rejected_mask
}
else:
chosen = torch.stack([item[0] for item in batch])
rejected = torch.stack([item[1] for item in batch])
return {'chosen': chosen, 'rejected': rejected}
collated = {}
instr_list = [item['instruction'] for item in batch if item.get('instruction') is not None]
if instr_list:
padded_instr = pad_sequence(instr_list, batch_first=True, padding_value=pad_token_id)
instr_mask = (padded_instr != pad_token_id).long()
collated['instruction'] = padded_instr
collated['instruction_mask'] = instr_mask
else:
collated['instruction'] = None
collated['instruction_mask'] = None
resp_list = [item['response'] for item in batch if item.get('response') is not None]
if resp_list:
padded_resp = pad_sequence(resp_list, batch_first=True, padding_value=pad_token_id)
resp_mask = (padded_resp != pad_token_id).long()
collated['response'] = padded_resp
collated['response_mask'] = resp_mask
else:
collated['response'] = None
collated['response_mask'] = None
modality_list = [item.get('modality_data') for item in batch if item.get('modality_data') is not None]
if modality_list and any(m is not None for m in modality_list):
images = [m.get('image') for m in modality_list if m and 'image' in m]
if images:
collated['modality_data'] = {'image': torch.stack(images)}
else:
collated['modality_data'] = None
else:
collated['modality_data'] = None
collated['task'] = [item.get('task', 'unknown') for item in batch]
return collated
return collate_fn_v2
def create_pretrain_dataloader(
mix_name: str = 'default',
tokenizer=None,
batch_size: int = 8,
num_workers: int = 4,
max_length: int = 2048,
max_samples: Optional[int] = None
):
if tokenizer.pad_token_id is None:
tokenizer.pad_token_id = tokenizer.eos_token_id
dataset = PreTrainDataset(
mix_name=mix_name,
tokenizer=tokenizer,
max_length=max_length,
streaming=True,
max_samples=max_samples
)
collate_fn = DynamicCollate(pad_token_id=tokenizer.pad_token_id)
return DataLoader(
dataset,
batch_size=batch_size,
num_workers=num_workers,
collate_fn=collate_fn,
pin_memory=True
)
def create_posttrain_dataloader(
mix_name: str = 'default',
tokenizer=None,
batch_size: int = 8,
num_workers: int = 4,
max_length: int = 2048,
max_samples: Optional[int] = None,
split: str = 'train',
shuffle: bool = True
):
if tokenizer.pad_token_id is None:
tokenizer.pad_token_id = tokenizer.eos_token_id
dataset = PostTrainDataset(
mix_name=mix_name,
tokenizer=tokenizer,
max_length=max_length,
max_samples=max_samples,
split=split
)
collate_fn = collate_fn_v2_factory(pad_token_id=tokenizer.pad_token_id)
return DataLoader(
dataset,
batch_size=batch_size,
shuffle=shuffle,
num_workers=num_workers,
collate_fn=collate_fn,
pin_memory=True,
drop_last=False
)
def create_preference_dataloader(
dataset_name: str = 'hh_rlhf',
tokenizer=None,
batch_size: int = 8,
num_workers: int = 4,
max_length: int = 1024,
max_samples: Optional[int] = None,
split: str = 'train',
shuffle: bool = True
):
dataset = PreferenceDataset(
dataset_name=dataset_name,
tokenizer=tokenizer,
max_length=max_length,
max_samples=max_samples,
split=split
)
return DataLoader(
dataset,
batch_size=batch_size,
shuffle=shuffle,
num_workers=num_workers,
collate_fn=collate_fn_v2_factory(pad_token_id=tokenizer.pad_token_id if tokenizer.pad_token_id is not None else tokenizer.eos_token_id),
pin_memory=True
)