| 38 | namespace { |
| 39 | |
| 40 | struct VariantValue { |
| 41 | string TypeName() const { return "TEST VariantValue"; } |
| 42 | static Status CPUZerosLikeFn(OpKernelContext* ctx, const VariantValue& v, |
| 43 | VariantValue* v_out) { |
| 44 | if (v.early_exit) { |
| 45 | return errors::InvalidArgument("early exit zeros_like!"); |
| 46 | } |
| 47 | v_out->value = 1; // CPU |
| 48 | return Status::OK(); |
| 49 | } |
| 50 | static Status GPUZerosLikeFn(OpKernelContext* ctx, const VariantValue& v, |
| 51 | VariantValue* v_out) { |
| 52 | if (v.early_exit) { |
| 53 | return errors::InvalidArgument("early exit zeros_like!"); |
| 54 | } |
| 55 | v_out->value = 2; // GPU |
| 56 | return Status::OK(); |
| 57 | } |
| 58 | static Status CPUAddFn(OpKernelContext* ctx, const VariantValue& a, |
| 59 | const VariantValue& b, VariantValue* out) { |
| 60 | if (a.early_exit) { |
| 61 | return errors::InvalidArgument("early exit add!"); |
| 62 | } |
| 63 | out->value = a.value + b.value; // CPU |
| 64 | return Status::OK(); |
| 65 | } |
| 66 | static Status GPUAddFn(OpKernelContext* ctx, const VariantValue& a, |
| 67 | const VariantValue& b, VariantValue* out) { |
| 68 | if (a.early_exit) { |
| 69 | return errors::InvalidArgument("early exit add!"); |
| 70 | } |
| 71 | out->value = -(a.value + b.value); // GPU |
| 72 | return Status::OK(); |
| 73 | } |
| 74 | static Status CPUToGPUCopyFn( |
| 75 | const VariantValue& from, VariantValue* to, |
| 76 | const std::function<Status(const Tensor&, Tensor*)>& copier) { |
| 77 | TF_RETURN_IF_ERROR(copier(Tensor(), nullptr)); |
| 78 | to->value = 0xdeadbeef; |
| 79 | return Status::OK(); |
| 80 | } |
| 81 | bool early_exit; |
| 82 | int value; |
| 83 | }; |
| 84 | |
| 85 | REGISTER_UNARY_VARIANT_DECODE_FUNCTION(VariantValue, "TEST VariantValue"); |
| 86 | |