"""
Data loaders for training and evaluation.

Provides efficient batching and chunking for both discrete tokens
and continuous vectors.
"""

import torch
from torch.utils.data import Dataset, DataLoader
from typing import List, Optional, Tuple, Iterator
import random


class TextDataset(Dataset):
    """
    Dataset for text data.

    Handles tokenization and sequence length management.
    """

    def __init__(
        self,
        texts: List[str],
        tokenizer,
        max_length: int = 512,
        stride: Optional[int] = None,
    ):
        self.texts = texts
        self.tokenizer = tokenizer
        self.max_length = max_length
        self.stride = stride or max_length

        # Pre-tokenize all texts
        self.tokenized = []
        for text in texts:
            token_ids = tokenizer.encode(text, add_special_tokens=True)
            self.tokenized.append(token_ids)

        # Create sliding windows for long sequences
        self.examples = []
        for token_ids in self.tokenized:
            if len(token_ids) <= max_length:
                self.examples.append(token_ids)
            else:
                # Sliding window
                for i in range(0, len(token_ids) - max_length + 1, self.stride):
                    window = token_ids[i:i + max_length]
                    self.examples.append(window)

    def __len__(self) -> int:
        return len(self.examples)

    def __getitem__(self, idx: int) -> torch.Tensor:
        token_ids = self.examples[idx]

        # Pad if needed
        if len(token_ids) < self.max_length:
            token_ids = token_ids + [self.tokenizer.pad_token_id] * (self.max_length - len(token_ids))

        return torch.tensor(token_ids, dtype=torch.long)


class ChunkedDataset(Dataset):
    """
    Dataset that returns chunks of K tokens for CALM training.

    Each example is a chunk of exactly K tokens, suitable for
    encoding into a single continuous vector.
    """

    def __init__(
        self,
        texts: List[str],
        tokenizer,
        chunk_size: int = 8,
    ):
        self.tokenizer = tokenizer
        self.chunk_size = chunk_size

        # Tokenize all texts
        all_token_ids = []
        for text in texts:
            token_ids = tokenizer.encode(text, add_special_tokens=False)
            all_token_ids.extend(token_ids)

        # Create chunks
        self.chunks = []
        for i in range(0, len(all_token_ids) - chunk_size + 1, chunk_size):
            chunk = all_token_ids[i:i + chunk_size]
            if len(chunk) == chunk_size:
                self.chunks.append(chunk)

    def __len__(self) -> int:
        return len(self.chunks)

    def __getitem__(self, idx: int) -> torch.Tensor:
        return torch.tensor(self.chunks[idx], dtype=torch.long)


class TextDataLoader:
    """
    Data loader for text sequences.

    Provides batching with proper padding and masking.
    """

    def __init__(
        self,
        dataset: Dataset,
        batch_size: int = 32,
        shuffle: bool = True,
        num_workers: int = 0,
        pin_memory: bool = True,
    ):
        self.dataset = dataset
        self.batch_size = batch_size
        self.shuffle = shuffle

        self.loader = DataLoader(
            dataset,
            batch_size=batch_size,
            shuffle=shuffle,
            num_workers=num_workers,
            pin_memory=pin_memory,
            collate_fn=self.collate_fn
        )

    def collate_fn(self, batch: List[torch.Tensor]) -> Tuple[torch.Tensor, torch.Tensor]:
        """
        Collate batch with padding.

        Returns:
            token_ids: [batch, max_len] padded token IDs
            mask: [batch, max_len] attention mask (1 for real tokens, 0 for padding)
        """
        # Stack tensors
        token_ids = torch.stack(batch, dim=0)

        # Create mask (1 for non-padding, 0 for padding)
        pad_token_id = self.dataset.tokenizer.pad_token_id
        mask = (token_ids != pad_token_id).long()

        return token_ids, mask

    def __iter__(self) -> Iterator[Tuple[torch.Tensor, torch.Tensor]]:
        return iter(self.loader)

    def __len__(self) -> int:
        return len(self.loader)


class ChunkedDataLoader:
    """
    Data loader for chunked data (CALM training).

    Returns batches of K-token chunks.
    """

    def __init__(
        self,
        dataset: ChunkedDataset,
        batch_size: int = 32,
        shuffle: bool = True,
        num_workers: int = 0,
        pin_memory: bool = True,
    ):
        self.dataset = dataset
        self.batch_size = batch_size
        self.shuffle = shuffle

        self.loader = DataLoader(
            dataset,
            batch_size=batch_size,
            shuffle=shuffle,
            num_workers=num_workers,
            pin_memory=pin_memory,
            collate_fn=self.collate_fn
        )

    def collate_fn(self, batch: List[torch.Tensor]) -> torch.Tensor:
        """
        Collate batch of chunks.

        Returns:
            chunks: [batch, chunk_size] token IDs
        """
        return torch.stack(batch, dim=0)

    def __iter__(self) -> Iterator[torch.Tensor]:
        return iter(self.loader)

    def __len__(self) -> int:
        return len(self.loader)


def load_text_file(file_path: str, encoding: str = 'utf-8') -> List[str]:
    """
    Load text file and split into lines or paragraphs.

    Args:
        file_path: Path to text file
        encoding: File encoding

    Returns:
        texts: List of text strings
    """
    with open(file_path, 'r', encoding=encoding) as f:
        content = f.read()

    # Split by double newline (paragraphs) or single newline (lines)
    if '\n\n' in content:
        texts = [p.strip() for p in content.split('\n\n') if p.strip()]
    else:
        texts = [line.strip() for line in content.split('\n') if line.strip()]

    return texts


def create_train_val_split(
    texts: List[str],
    val_ratio: float = 0.1,
    seed: int = 42
) -> Tuple[List[str], List[str]]:
    """
    Split texts into train and validation sets.

    Args:
        texts: List of text strings
        val_ratio: Validation set ratio
        seed: Random seed

    Returns:
        train_texts: Training texts
        val_texts: Validation texts
    """
    random.seed(seed)
    texts_copy = texts.copy()
    random.shuffle(texts_copy)

    split_idx = int(len(texts_copy) * (1 - val_ratio))
    train_texts = texts_copy[:split_idx]
    val_texts = texts_copy[split_idx:]

    return train_texts, val_texts
