import os
import pkgutil
from copy import deepcopy
from typing import Optional

from torch import nn as nn

from timm.layers import Conv2dSame, BatchNormAct2d, Linear

__all__ = ['extract_layer', 'set_layer', 'adapt_model_from_string', 'adapt_model_from_file']


def extract_layer(model, layer):
    """Extract a layer from a model using dot-separated path.

    Args:
        model: PyTorch model.
        layer: Dot-separated layer path (e.g., 'layer1.0.conv1').

    Returns:
        Extracted module.
    """
    layer = layer.split('.')
    module = model
    if hasattr(model, 'module') and layer[0] != 'module':
        module = model.module
    if not hasattr(model, 'module') and layer[0] == 'module':
        layer = layer[1:]
    for l in layer:
        if hasattr(module, l):
            if not l.isdigit():
                module = getattr(module, l)
            else:
                module = module[int(l)]
        else:
            return module
    return module


def set_layer(model, layer, val):
    """Set a layer in a model using dot-separated path.

    Args:
        model: PyTorch model.
        layer: Dot-separated layer path.
        val: New value for the layer.
    """
    layer = layer.split('.')
    module = model
    if hasattr(model, 'module') and layer[0] != 'module':
        module = model.module
    lst_index = 0
    module2 = module
    for l in layer:
        if hasattr(module2, l):
            if not l.isdigit():
                module2 = getattr(module2, l)
            else:
                module2 = module2[int(l)]
            lst_index += 1
    lst_index -= 1
    for l in layer[:lst_index]:
        if not l.isdigit():
            module = getattr(module, l)
        else:
            module = module[int(l)]
    l = layer[lst_index]
    setattr(module, l, val)


def adapt_model_from_string(parent_module, model_string):
    """Adapt a model to pruned structure from string specification.

    Args:
        parent_module: Original model to adapt.
        model_string: String containing layer shapes for pruned model.

    Returns:
        Adapted model with pruned layer dimensions.
    """
    separator = '***'
    state_dict = {}
    lst_shape = model_string.split(separator)
    for k in lst_shape:
        k = k.split(':')
        key = k[0]
        shape = k[1][1:-1].split(',')
        if shape[0] != '':
            state_dict[key] = [int(i) for i in shape]

    # Extract device and dtype from the parent module
    device = next(parent_module.parameters()).device
    dtype = next(parent_module.parameters()).dtype
    dd = {'device': device, 'dtype': dtype}

    input_convs = getattr(parent_module, 'pretrained_cfg', {}).get('first_conv', ())
    if isinstance(input_convs, str):
        input_convs = (input_convs,)

    new_module = deepcopy(parent_module)
    for n, m in parent_module.named_modules():
        old_module = extract_layer(parent_module, n)
        if isinstance(old_module, nn.Conv2d) or isinstance(old_module, Conv2dSame):
            if isinstance(old_module, Conv2dSame):
                conv = Conv2dSame
            else:
                conv = nn.Conv2d
            s = state_dict[n + '.weight']
            in_channels = old_module.in_channels if n in input_convs else s[1]
            out_channels = s[0]
            g = 1
            if old_module.groups > 1:
                in_channels = out_channels
                g = in_channels
            new_conv = conv(
                in_channels=in_channels,
                out_channels=out_channels,
                kernel_size=old_module.kernel_size,
                bias=old_module.bias is not None,
                padding=old_module.padding,
                dilation=old_module.dilation,
                groups=g,
                stride=old_module.stride,
                **dd,
            )
            set_layer(new_module, n, new_conv)
        elif isinstance(old_module, BatchNormAct2d):
            new_bn = BatchNormAct2d(
                state_dict[n + '.weight'][0],
                eps=old_module.eps,
                momentum=old_module.momentum,
                affine=old_module.affine,
                track_running_stats=True,
                **dd,
            )
            new_bn.drop = old_module.drop
            new_bn.act = old_module.act
            set_layer(new_module, n, new_bn)
        elif isinstance(old_module, nn.BatchNorm2d):
            new_bn = nn.BatchNorm2d(
                num_features=state_dict[n + '.weight'][0],
                eps=old_module.eps,
                momentum=old_module.momentum,
                affine=old_module.affine,
                track_running_stats=True,
                **dd,
            )
            set_layer(new_module, n, new_bn)
        elif isinstance(old_module, nn.Linear):
            # FIXME extra checks to ensure this is actually the FC classifier layer and not a diff Linear layer?
            num_features = state_dict[n + '.weight'][1]
            new_fc = Linear(
                in_features=num_features,
                out_features=old_module.out_features,
                bias=old_module.bias is not None,
                **dd,
            )
            set_layer(new_module, n, new_fc)
            if hasattr(new_module, 'num_features'):
                if getattr(new_module, 'head_hidden_size', 0) == new_module.num_features:
                    new_module.head_hidden_size = num_features
                new_module.num_features = num_features

    new_module.eval()
    parent_module.eval()

    # The rebuilt layers have new widths, but feature_info still carries the unpruned channel counts.
    _adapt_feature_info(new_module)

    return new_module


def _module_out_chs(module: nn.Module) -> Optional[int]:
    """Return the output width of a conv / norm / linear layer, or None for any other module."""
    if isinstance(module, nn.Conv2d):
        return module.out_channels
    if isinstance(module, nn.BatchNorm2d):
        return module.num_features
    if isinstance(module, nn.Linear):
        return module.out_features
    return None


def _adapt_feature_info(module: nn.Module) -> None:
    """Update ``feature_info`` channel counts in-place to match a pruned module tree.

    Each feature entry names the module producing that feature. After pruning, its output width is the
    width of the last conv / norm / linear layer registered within it. A parameter-less feature module,
    such as a stem activation, takes the width of the last such layer registered before it.

    Args:
        module: Pruned model whose ``feature_info`` (a list of dicts or a ``FeatureInfo``) is updated.
    """
    feature_info = getattr(module, 'feature_info', None)
    infos = getattr(feature_info, 'info', feature_info)
    if not isinstance(infos, (list, tuple)):
        return
    entries = [(i, info['module']) for i, info in enumerate(infos) if isinstance(info, dict) and info.get('module')]
    before = {}
    within = {}
    current = None
    for name, m in module.named_modules():
        for i, prefix in entries:
            if name == prefix:
                before[i] = current
        out_chs = _module_out_chs(m)
        if out_chs is None:
            continue
        current = out_chs
        for i, prefix in entries:
            if name == prefix or name.startswith(prefix + '.'):
                within[i] = out_chs
    for i, _ in entries:
        num_chs = within.get(i, before.get(i))
        if num_chs is not None:
            infos[i]['num_chs'] = num_chs


def adapt_model_from_file(parent_module, model_variant):
    """Adapt a model to pruned structure from file specification.

    Args:
        parent_module: Original model to adapt.
        model_variant: Name of pruned model variant file.

    Returns:
        Adapted model with pruned layer dimensions.
    """
    adapt_data = pkgutil.get_data(__name__, os.path.join('_pruned', model_variant + '.txt'))
    return adapt_model_from_string(parent_module, adapt_data.decode('utf-8').strip())
