| 352 | } |
| 353 | |
| 354 | ValueRefList ScalarTransformation::apply_transformation( |
| 355 | const Operator& op, Span<ValueRef> inputs) { |
| 356 | if (auto* get_attr = op.as<GetAttr>()) { |
| 357 | // fastpath for GetAttr |
| 358 | return apply_get_attr(*get_attr, inputs); |
| 359 | } else if (auto* apply_op = op.as<ApplyOp>()) { |
| 360 | if (apply_op->op().same_type<FastpathCopy>()) { |
| 361 | return inputs[0]; |
| 362 | } |
| 363 | } |
| 364 | size_t nr_inputs = inputs.size(); |
| 365 | ValueRefList unwrapped_inputs(nr_inputs); |
| 366 | SmallVector<bool> inputs_mask(nr_inputs); |
| 367 | for (size_t i = 0; i < inputs.size(); ++i) { |
| 368 | if (auto&& scalar_value = inputs[i].as_ref(m_value_type)) { |
| 369 | unwrapped_inputs[i] = scalar_value->value(); |
| 370 | inputs_mask[i] = true; |
| 371 | } else { |
| 372 | unwrapped_inputs[i] = inputs[i]; |
| 373 | inputs_mask[i] = false; |
| 374 | } |
| 375 | } |
| 376 | auto fallback = [&] { return imperative::apply(op, unwrapped_inputs); }; |
| 377 | if (auto apply_op = op.as<ApplyOp>()) { |
| 378 | auto iter = scalar_rules.find(apply_op->op().dyn_typeinfo()); |
| 379 | if (iter != scalar_rules.end()) { |
| 380 | return iter->second( |
| 381 | apply_op->op(), unwrapped_inputs, inputs_mask, m_value_type); |
| 382 | } else { |
| 383 | // TODO: repeat op |
| 384 | return fallback(); |
| 385 | } |
| 386 | } else if (auto* create_tensor = op.as<CreateTensor>()) { |
| 387 | if (create_tensor->shape().is_scalar()) { |
| 388 | ValueShape scalar_shape = {1}; |
| 389 | CreateTensor scalar_op( |
| 390 | create_tensor->kind(), create_tensor->device(), |
| 391 | create_tensor->dtype(), scalar_shape); |
| 392 | return {m_value_type.make(imperative::apply(scalar_op, inputs)[0])}; |
| 393 | } else { |
| 394 | return imperative::apply(op, inputs); |
| 395 | } |
| 396 | } else if (op.as<IsScalar>()) { |
| 397 | mgb_assert(nr_inputs == 1); |
| 398 | return {BoolValue::make(inputs_mask[0])}; |
| 399 | } else if (op.is<Operator::IdentityLike>()) { |
| 400 | mgb_assert(nr_inputs == 1); |
| 401 | bool is_scalar = inputs_mask[0]; |
| 402 | auto outputs = fallback(); |
| 403 | if (is_scalar) { |
| 404 | outputs[0] = m_value_type.make(outputs[0]); |
| 405 | } |
| 406 | return outputs; |
| 407 | } else { |
| 408 | return fallback(); |
| 409 | } |
| 410 | }; |
| 411 | |