apple / apple/coremltools

Mismatch in tensor shapes when converting LogBinomial

Open
#2,020 1 comment 0 reactions 0 assignees View on GitHub
bug triaged
Dominant language
Python
Stars
5.4k
Forks
850
Avg merge
4d 5h
Merged PRs (30d)
10

Description

I'm trying to convert a model that contains a `LogBinomial` module and I'm getting a mismatch on some tensor shapes when converting, I was hoping you would have more insight in solving this issue. Thank you.

## Stack Trace
```
Converting PyTorch Frontend ==> MIL Ops: 98%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████▊ | 49/50 [00:00<00:00, 9883.19 ops/s]
Running MIL frontend_pytorch pipeline: 100%|██████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 5/5 [00:00<00:00, 7576.42 passes/s]
Running MIL default pipeline: 55%|██████████████████████████████████████████████████████████████████████████████████████████ | 36/66 [00:00<00:00, 3723.67 passes/s]
Traceback (most recent call last):
File "/Users/guiye/Downloads/omnidata-main/omnidata_tools/ZoeDepth/logbinomial.py", line 86, in
ct_model = ct.convert(
File "/Users/guiye/miniforge3/envs/zoe/lib/python3.9/site-packages/coremltools/converters/_converters_entry.py", line 551, in convert
mlmodel = mil_convert(
File "/Users/guiye/miniforge3/envs/zoe/lib/python3.9/site-packages/coremltools/converters/mil/converter.py", line 188, in mil_convert
return _mil_convert(model, convert_from, convert_to, ConverterRegistry, MLModel, compute_units, **kwargs)
File "/Users/guiye/miniforge3/envs/zoe/lib/python3.9/site-packages/coremltools/converters/mil/converter.py", line 212, in _mil_convert
proto, mil_program = mil_convert_to_proto(
File "/Users/guiye/miniforge3/envs/zoe/lib/python3.9/site-packages/coremltools/converters/mil/converter.py", line 289, in mil_convert_to_proto
PassPipelineManager.apply_pipeline(prog, main_pipeline)
File "/Users/guiye/miniforge3/envs/zoe/lib/python3.9/site-packages/coremltools/converters/mil/mil/passes/pass_pipeline.py", line 448, in apply_pipeline
graph_pass(prog)
File "/Users/guiye/miniforge3/envs/zoe/lib/python3.9/site-packages/coremltools/converters/mil/mil/passes/graph_pass.py", line 51, in __call__
self.apply(prog)
File "/Users/guiye/miniforge3/envs/zoe/lib/python3.9/site-packages/coremltools/converters/mil/mil/passes/defs/optimize_elementwise_binary.py", line 255, in apply
block_changed = self._fuse_elementwise_to_batchnorm_block(f)
File "/Users/guiye/miniforge3/envs/zoe/lib/python3.9/site-packages/coremltools/converters/mil/mil/passes/helper.py", line 60, in wrapper
return func(*args)
File "/Users/guiye/miniforge3/envs/zoe/lib/python3.9/site-packages/coremltools/converters/mil/mil/passes/defs/optimize_elementwise_binary.py", line 347, in _fuse_elementwise_to_batchnorm_block
fusion_status = self._try_to_transform(op, add_op, block)
File "/Users/guiye/miniforge3/envs/zoe/lib/python3.9/site-packages/coremltools/converters/mil/mil/passes/defs/optimize_elementwise_binary.py", line 326, in _try_to_transform
add_op.enclosing_block.replace_uses_of_var_after_op(
File "/Users/guiye/miniforge3/envs/zoe/lib/python3.9/site-packages/coremltools/converters/mil/mil/block.py", line 630, in replace_uses_of_var_after_op
num_ops_affected = self._replace_var(
File "/Users/guiye/miniforge3/envs/zoe/lib/python3.9/site-packages/coremltools/converters/mil/mil/block.py", line 403, in _replace_var
op.set_inputs(no_check_var_types=no_check_var_types,
File "/Users/guiye/miniforge3/envs/zoe/lib/python3.9/site-packages/coremltools/converters/mil/mil/operation.py", line 225, in set_inputs
self._validate_and_set_inputs(input_kvs, no_check_var_types=no_check_var_types)
File "/Users/guiye/miniforge3/envs/zoe/lib/python3.9/site-packages/coremltools/converters/mil/mil/operation.py", line 510, in _validate_and_set_inputs
check_and_detach(
File "/Users/guiye/miniforge3/envs/zoe/lib/python3.9/site-packages/coremltools/converters/mil/mil/operation.py", line 493, in check_and_detach
raise ValueError(
ValueError: New var type `.tensor'>` not a subtype of existing var type `.tensor'>`.
```

## To Reproduce
``` python
import torch
import torch.nn as nn
import numpy as np
import coremltools as ct

def log_binom(n, k, eps=1e-7):
""" log(nCk) using stirling approximation """
n = n + eps
k = k + eps
return n * torch.log(n) - k * torch.log(k) - (n-k) * torch.log(n-k+eps)

class LogBinomial(nn.Module):
def __init__(self, n_classes=256, act=torch.softmax):
"""Compute log binomial distribution for n_classes

Args:
n_classes (int, optional): number of output classes. Defaults to 256.
"""
super().__init__()
self.K = n_classes
self.act = act
self.register_buffer('k_idx', torch.arange(
0, n_classes).view(1, -1, 1, 1))
self.register_buffer('K_minus_1', torch.Tensor(
[self.K-1]).view(1, -1, 1, 1))

def forward(self, x, t=1., eps=1e-4):
"""Compute log binomial distribution for x

Args:
x (torch.Tensor - NCHW): probabilities
t (float, torch.Tensor - NCHW, optional): Temperature of distribution. Defaults to 1..
eps (float, optional): Small number for numerical stability. Defaults to 1e-4.

Returns:
torch.Tensor -NCHW: log binomial distribution logbinomial(p;t)
"""
if x.ndim == 3:
x = x.unsqueeze(1) # make it nchw

one_minus_x = torch.clamp(1 - x, eps, 1)
x = torch.clamp(x, eps, 1)
y = log_binom(self.K_minus_1, self.k_idx) + self.k_idx * \
torch.log(x) + (self.K - 1 - self.k_idx) * torch.log(one_minus_x)
return self.act(y/t, dim=1)

model = LogBinomial(64).eval()
x_tensor = torch.rand(1, 512, 512)
traced_model = torch.jit.trace(model, x_tensor)

ct_model = ct.convert(
traced_model,
convert_to="mlprogram",
minimum_deployment_target=ct.target.macOS13,
inputs=[ct.TensorType(name='x', shape=x_tensor.shape)],
outputs=[ct.TensorType(name='out', dtype=np.float32)],
)

ct_model.save("LogBinomial.mlpackage")
```

## System environment (please complete the following information):
- coremltools version: 7.0
- PyTorch: 2.0.0
- OS: macOS 14.0

## Additional context
`LogBinomial` module is taken directly from the [ZoeDepth](https://github.com/isl-org/ZoeDepth) codebase.

Contributor guide

Open the contributing guide

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.