(mask, x, y)
| 217 | |
| 218 | |
| 219 | def where(mask, x, y): |
| 220 | assert isinstance(mask, HLOTensor), f"mask must be HLOTensor, get {type(mask)}" |
| 221 | x = x if isinstance(x, HLOTensor) else HLOTensor(x) |
| 222 | y = y if isinstance(y, HLOTensor) else HLOTensor(y) |
| 223 | |
| 224 | return mask * x + (np.array(1.0).astype(x.dtype) - mask) * y |
| 225 | |
| 226 | |
| 227 | def where_grad(dout, mask): |
no test coverage detected