File size: 10,090 Bytes
31dc8dc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
#!/usr/bin/env python
"""Evaluation script for trained ELF models: loads a checkpoint and generates text samples."""

import argparse
import contextlib
import copy
import logging
import os
import sys

# Initialize JAX distributed BEFORE importing other JAX modules
import jax
try:
    jax.distributed.initialize()
except (RuntimeError, ValueError):
    pass  # Single-host run, or already initialized.

# Ensure repo root on sys.path so imports work when run as a script
REPO_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
if REPO_ROOT not in sys.path:
    sys.path.insert(0, REPO_ROOT)

import jax.numpy as jnp
import optax
from flax import jax_utils
from transformers import AutoTokenizer

from modules.t5_encoder import get_encoder
from modules.model import ELF_models
from utils.logging_utils import log_for_0
from utils.checkpoint_utils import load_encoder_checkpoint, load_checkpoint
from utils.train_utils import TrainState
from utils.data_utils import load_jsonl_dataset, load_dataset_split, get_pad_token_id
from generation import test_generation_uncond, test_generation_cond
from configs.config import load_config_from_yaml, apply_config_overrides, load_sampling_configs

logging.basicConfig(
    format="%(levelname)s - %(name)s - %(message)s",
    handlers=[logging.StreamHandler(sys.stdout)],
    level=logging.INFO, force=True,
)
logger = logging.getLogger(__name__)


def parse_args():
    parser = argparse.ArgumentParser(description="Evaluate trained ELF model by generating text samples")
    parser.add_argument("--config", type=str, required=True, help="Path to configuration YAML file")
    parser.add_argument(
        "--config_override", action="append", default=[],
        help="Override config values (field_name=value). Can be specified multiple times.",
    )
    parser.add_argument("--seed", type=int, default=42, help="Random seed (used when --seeds is not specified)")
    parser.add_argument(
        "--seeds", type=str, default=None,
        help="Comma-separated list of seeds to evaluate (e.g. '42,123,456'). Overrides --seed.",
    )
    parser.add_argument(
        "--checkpoint_path", type=str, required=True,
        help="Path to checkpoint file (e.g. outputs/elf_b-owt/checkpoint_19000) or HF repo id.",
    )
    parser.add_argument(
        "--use_cpu", action="store_true",
        help="Host model init, train state template, and encoder/state replication on CPU",
    )
    return parser.parse_args()


def main():
    args = parse_args()

    log_for_0("Loading configuration...")
    config = load_config_from_yaml(args.config)
    if args.config_override:
        config = apply_config_overrides(config, args.config_override)
        log_for_0(f"Applied {len(args.config_override)} config override(s)")

    num_devices = jax.device_count()
    num_local_devices = jax.local_device_count()
    num_hosts = jax.process_count()
    cpu_device = jax.local_devices(backend="cpu")[0] if args.use_cpu else None

    def cpu_ctx():
        return jax.default_device(cpu_device) if args.use_cpu else contextlib.nullcontext()

    if config.global_batch_size is not None:
        log_for_0(f"Using global batch size for evaluation: {config.global_batch_size}")
        total_batch_size = config.global_batch_size
        local_batch_size = total_batch_size // num_hosts
        config.batch_size = local_batch_size
    elif config.batch_size is not None:
        log_for_0(f"Using batch size per device: {config.batch_size}")
        total_batch_size = config.batch_size * num_devices
        local_batch_size = config.batch_size * num_local_devices
        config.global_batch_size = total_batch_size
    else:
        raise ValueError("Either global_batch_size or batch_size must be specified")

    log_for_0(f"Config loaded from {args.config}")
    log_for_0(f"Model: {config.model}")
    log_for_0(f"Encoder Model: {config.encoder_model_name}")
    log_for_0(f"Encoder Checkpoint: {config.encoder_checkpoint}")
    log_for_0(f"Max length: {config.max_length}")
    log_for_0(f"Max input length: {config.max_input_length}")
    log_for_0(f"Num samples: {config.num_samples}")
    log_for_0(f"Sampling configs: {len(config.sampling_configs)} config(s)")

    seed_list = [int(s.strip()) for s in args.seeds.split(",")] if args.seeds is not None else [args.seed]
    log_for_0(f"Seeds to evaluate: {seed_list}")

    rng = jax.random.PRNGKey(config.seed)

    log_for_0("Loading tokenizer...")
    tokenizer = AutoTokenizer.from_pretrained(config.tokenizer_name or config.encoder_model_name)
    pad_token_id = get_pad_token_id(tokenizer, config.pad_token)
    log_for_0(f"Using {'EOS' if config.pad_token == 'eos' else 'PAD'} token for padding: {pad_token_id}")

    eval_dataset = None
    if config.eval_data_path is not None:
        log_for_0("Loading dataset for conditional generation...")
        if config.eval_data_path.endswith(".jsonl"):
            eval_dataset = load_jsonl_dataset(
                config.eval_data_path, tokenizer,
                input_key="input",
                output_key="output",
            )
        else:
            eval_dataset = load_dataset_split(config.eval_data_path)
        log_for_0(f"Eval dataset size: {len(eval_dataset)}")

    # ============================================
    # Load Encoder (frozen)
    # ============================================
    log_for_0(f"Loading Encoder config: {config.encoder_model_name}...")
    encoder_config, encoder_model, _ = get_encoder(config.encoder_model_name, jnp.float32)
    encoder_params = load_encoder_checkpoint(config.encoder_checkpoint)
    log_for_0("encoder weights loaded.")

    # Multi-device eval passes encoder params directly into pmap, so replicate them
    # across local accelerator devices.
    encoder_params = jax_utils.replicate(encoder_params)
    log_for_0(f"Encoder d_model: {encoder_config.d_model}")

    # ============================================
    # Create ELF Model
    # ============================================
    log_for_0(f"Creating {config.model} model...")
    rng, init_rng, dropout_rng = jax.random.split(rng, 3)
    max_length = config.max_length

    with cpu_ctx():
        # 2x dim if self_cond_prob > 0 to initialize self_cond_proj layer
        _text_enc_dim = encoder_config.d_model
        input_dim = 2 * _text_enc_dim if config.self_cond_prob > 0 else _text_enc_dim
        dummy_x = jnp.ones((1, max_length, input_dim))
        dummy_t = jnp.ones((1,))
        dummy_self_cond_cfg_scale = jnp.ones((1,)) if config.num_self_cond_cfg_tokens > 0 else None
        log_for_0(f"Dummy x shape: {dummy_x.shape}")
        log_for_0(f"Dummy t shape: {dummy_t.shape}")

    vocab_size = tokenizer.vocab_size
    model = ELF_models[config.model](
        text_encoder_dim=encoder_config.d_model,
        max_length=max_length,
        attn_drop=config.attn_dropout,
        proj_drop=config.proj_dropout,
        num_time_tokens=config.num_time_tokens,
        num_self_cond_cfg_tokens=config.num_self_cond_cfg_tokens,
        vocab_size=vocab_size,
        num_model_mode_tokens=config.num_model_mode_tokens,
        bottleneck_dim=config.bottleneck_dim,
    )

    log_for_0("Initializing ELF model...")
    init_args = dict(
        x=dummy_x, t=dummy_t, deterministic=True,
        self_cond_cfg_scale=dummy_self_cond_cfg_scale,
    )
    with cpu_ctx():
        elf_params = model.init(init_rng, **init_args)
        log_for_0("\n" + model.tabulate(init_rng, **init_args))
        log_for_0("ELF initialization complete")

    total_params = sum(x.size for x in jax.tree_util.tree_leaves(elf_params))
    log_for_0(f"ELF parameters: {total_params:,}")

    # ============================================
    # Create Train State Template
    # ============================================
    optimizer = optax.adamw(learning_rate=1e-4)
    with cpu_ctx():
        state = TrainState.create(
            apply_fn=model.apply,
            params=elf_params["params"],
            tx=optimizer,
            dropout_rng=dropout_rng,
            ema_params1=copy.deepcopy(elf_params["params"]),
        )

    # ============================================
    # Determine checkpoints to evaluate
    # ============================================
    if config.sampling_configs_path:
        config.sampling_configs = load_sampling_configs(config.sampling_configs_path)

    log_for_0(f"Loading checkpoint from: {args.checkpoint_path}")
    state, _ = load_checkpoint(args.checkpoint_path, state)
    state_replicated = jax_utils.replicate(state)

    for seed_idx, seed_val in enumerate(seed_list):
        if len(seed_list) > 1:
            log_for_0(f"\n{'#' * 70}")
            log_for_0(f"Seed {seed_idx + 1}/{len(seed_list)}: {seed_val}")
            log_for_0(f"{'#' * 70}")

        seed_rng = jax.random.PRNGKey(seed_val)

        original_output_dir = config.output_dir
        if len(seed_list) > 1:
            config.output_dir = os.path.join(original_output_dir, f"seed_{seed_val}")

        for sc_idx, sc in enumerate(config.sampling_configs):
            if len(config.sampling_configs) > 1:
                log_for_0(f"\n--- Sampling config {sc_idx + 1}/{len(config.sampling_configs)} ---")
            seed_rng, sample_rng = jax.random.split(seed_rng)
            common_kwargs = dict(
                state=state_replicated,
                tokenizer=tokenizer,
                rng=sample_rng,
                config=config,
                sampling_config=sc,
                batch_size=local_batch_size,
                num_samples=config.num_samples,
            )
            if eval_dataset is None:
                test_generation_uncond(**common_kwargs)
            else:
                test_generation_cond(
                    **common_kwargs,
                    encoder_params=encoder_params,
                    encoder_apply_fn=encoder_model.apply,
                    dataset=eval_dataset,
                )

        config.output_dir = original_output_dir

    log_for_0("\nEvaluation complete!")


if __name__ == "__main__":
    main()