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: """获取响应(兼容 / 标签)""" 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'([\\s\\S]*?)', last, re.IGNORECASE) if m: return m.group(1).strip() return re.sub(r'[\\s\\S]*?', '', 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 )