#!/usr/bin/env python3
"""Lab 08: local, learner-transcribed JSON -> calculated payloads, never RAM measurements."""
import argparse
import json
from pathlib import Path


def integer(value, name, minimum=1):
    if type(value) is not int or value < minimum:
        raise ValueError(f'{name} must be an integer >= {minimum}')
    return value


def payload(byte_count):
    return {'bytes': byte_count, 'MiB': byte_count / 2**20,
            'GiB': byte_count / 2**30}


def unique_keys(pairs):
    result = {}
    for key, value in pairs:
        if key in result:
            raise ValueError(f'Duplicate JSON field: {key}')
        result[key] = value
    return result


def calculate(spec):
    if not isinstance(spec, dict):
        raise ValueError('Input must be a JSON object')
    required = ('residual_width', 'layers', 'query_heads', 'kv_heads', 'head_dim',
                'batch', 'sequence_length', 'cache_bytes_per_scalar')
    values = {key: integer(spec.get(key), key) for key in required}
    local = integer(spec.get('local_layers'), 'local_layers', 0)
    global_ = integer(spec.get('global_layers'), 'global_layers', 0)
    if local + global_ != values['layers']:
        raise ValueError('Local and global layer counts must sum to layers')
    if values['query_heads'] % values['kv_heads']:
        raise ValueError('Query heads must divide into equal KV groups')
    if values['cache_bytes_per_scalar'] not in (1, 2, 4, 8):
        raise ValueError('Explicit cache scalar bytes must be 1, 2, 4 or 8')
    if ('unique_parameters' in spec) != ('weight_bits_per_scalar' in spec):
        raise ValueError('Supply parameter count and weight bits together, or omit both')
    weight_bytes = None
    if 'unique_parameters' in spec:
        parameters = integer(spec['unique_parameters'], 'unique_parameters')
        precision = integer(spec['weight_bits_per_scalar'], 'weight_bits_per_scalar')
        if precision not in (1, 2, 4, 8, 16, 32, 64):
            raise ValueError('Explicit weight bits must be 1, 2, 4, 8, 16, 32 or 64')
        bits = parameters * precision
        if bits % 8:
            raise ValueError('Packed weight total is not byte-aligned; no implicit rounding')
        weight_bytes = bits // 8
    window = integer(spec.get('local_window'), 'local_window') if local else 0
    source = spec.get('source')
    if (not isinstance(source, dict) or
        not isinstance(source.get('url'), str) or not source['url'].startswith('https://') or
        not isinstance(source.get('revision'), str) or len(source['revision']) != 40 or
        any(c not in '0123456789abcdef' for c in source['revision'])):
        raise ValueError('Provide an HTTPS source URL and full lowercase 40-hex revision')
    factor = (2 * values['batch'] * values['kv_heads'] * values['head_dim'] *
              values['cache_bytes_per_scalar'])
    full = factor * values['layers'] * values['sequence_length']
    bounded = factor * (global_ * values['sequence_length'] +
                        local * min(values['sequence_length'], window))
    d = values['residual_width']
    return {'kind': 'calculated payload; not measured RAM or VRAM',
            'input_kind': 'learner-created normalized transcription; not publisher config',
            'source': source, 'assumptions': spec,
            'excluded': ['quantization metadata', 'padding', 'mixed-precision exceptions',
                         'activations', 'workspace', 'allocator/runtime overhead',
                         'replication', 'cross-sequence cache sharing'],
            'logical_projection_shapes': {
                'Q': [d, values['query_heads'] * values['head_dim']],
                'K_and_V': [d, values['kv_heads'] * values['head_dim']],
                'O': [values['query_heads'] * values['head_dim'], d]},
            'full_length_KV': payload(full), 'bounded_local_KV': payload(bounded),
            'ideal_weight_payload': payload(weight_bytes) if weight_bytes is not None else None,
            'weights_plus_full_KV': payload(weight_bytes + full) if weight_bytes is not None else None}


def gpt2_parameters(vocabulary, width, positions, layers):
    for value, name in zip((vocabulary, width, positions, layers),
                           ('vocabulary', 'width', 'positions', 'layers')):
        integer(value, name)
    return vocabulary * width + positions * width + layers * (12 * width**2 + 13 * width) + 2 * width


def quantized_groups(scalars, bits, group_size, scale_bytes):
    for value, name in zip((scalars, bits, group_size, scale_bytes),
                           ('scalars', 'bits', 'group_size', 'scale_bytes')):
        integer(value, name)
    if scalars % group_size or scalars * bits % 8:
        raise ValueError('This simplified layout requires full groups and byte-aligned values')
    values = scalars * bits // 8
    scales = scalars // group_size * scale_bytes
    return {'value_bytes': values, 'scale_bytes': scales, 'total_bytes': values + scales}


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('input', type=Path, help='Your normalized local JSON worksheet')
    args = parser.parse_args()
    print(json.dumps(calculate(json.loads(args.input.read_text(encoding='utf-8'),
                                         object_pairs_hook=unique_keys)), indent=2))


if __name__ == '__main__':
    main()
