| 103 | } |
| 104 | |
| 105 | void RegisterBitcastOp() { |
| 106 | TF_Status* status = TF_NewStatus(); |
| 107 | |
| 108 | TF_OpDefinitionBuilder* op_builder = TF_NewOpDefinitionBuilder("Bitcast"); |
| 109 | TF_OpDefinitionBuilderAddInput(op_builder, "input: T"); |
| 110 | TF_OpDefinitionBuilderAddOutput(op_builder, "output: type"); |
| 111 | TF_OpDefinitionBuilderAddAttr( |
| 112 | op_builder, |
| 113 | "T: {bfloat16, half, float, double, int64, int32, uint8, uint16, " |
| 114 | "uint32, uint64, int8, int16, complex64, complex128, qint8, quint8, " |
| 115 | "qint16, quint16, qint32}"); |
| 116 | TF_OpDefinitionBuilderAddAttr( |
| 117 | op_builder, |
| 118 | "type: {bfloat16, half, float, double, int64, int32, uint8, uint16, " |
| 119 | "uint32, uint64, int8, int16, complex64, complex128, qint8, quint8, " |
| 120 | "qint16, quint16, qint32}"); |
| 121 | TF_OpDefinitionBuilderSetShapeInferenceFunction(op_builder, |
| 122 | &bitcast_shape_inference_fn); |
| 123 | |
| 124 | TF_RegisterOpDefinition(op_builder, status); |
| 125 | CHECK_EQ(TF_GetCode(status), TF_OK) |
| 126 | << "Bitcast op registration failed: " << TF_Message(status); |
| 127 | TF_DeleteStatus(status); |
| 128 | } |
| 129 | |
| 130 | static bool IsBitcastOpRegistered = []() { |
| 131 | if (SHOULD_REGISTER_OP("Bitcast")) { |
no test coverage detected