@@ -573,6 +573,7 @@ struct AggregateArgType {
573573 X(Int) \
574574 X(Float) \
575575 X(None_) \
576+ X(String) \
576577 X(Enum) \
577578 X(NativeDType) \
578579 X(ForeignDType)
@@ -651,6 +652,7 @@ enum class PythonArgKind : uint8_t {
651652 ConstantInt,
652653 ConstantFloat,
653654 ConstantNone,
655+ ConstantString,
654656 IdentityConstant,
655657 ForeignDTypeConstant,
656658 // A torch.Tensor that we can access via torch._C._to_dlpack
@@ -675,6 +677,7 @@ static inline PythonArgKind constant_kind_as_arg_kind(ConstantKind kind) {
675677 case ConstantKind::Int: return PythonArgKind::ConstantInt;
676678 case ConstantKind::Float: return PythonArgKind::ConstantFloat;
677679 case ConstantKind::None_: return PythonArgKind::ConstantNone;
680+ case ConstantKind::String: return PythonArgKind::ConstantString;
678681 case ConstantKind::Enum: return PythonArgKind::IdentityConstant;
679682 case ConstantKind::NativeDType: return PythonArgKind::IdentityConstant;
680683 case ConstantKind::ForeignDType: return PythonArgKind::ForeignDTypeConstant;
@@ -689,6 +692,7 @@ static ParameterKind::Category param_category_from_pyarg_kind(PythonArgKind k) {
689692 case PythonArgKind::ConstantInt: return ParameterKind::ConstantInt;
690693 case PythonArgKind::ConstantFloat: return ParameterKind::ConstantFloat;
691694 case PythonArgKind::ConstantNone: return ParameterKind::ConstantNone;
695+ case PythonArgKind::ConstantString: return ParameterKind::IdentityConstant;
692696 case PythonArgKind::IdentityConstant: return ParameterKind::IdentityConstant;
693697 case PythonArgKind::ForeignDTypeConstant: return ParameterKind::IdentityConstant;
694698 case PythonArgKind::TorchTensorDlpack: return ParameterKind::Array;
@@ -915,6 +919,9 @@ static std::optional<ConstantKind> classify_constant(PyObject* obj, bool kernel_
915919 if (obj == Py_None)
916920 return ConstantKind::None_;
917921
922+ if (PyUnicode_CheckExact (obj))
923+ return ConstantKind::String;
924+
918925 if (PyObject_TypeCheck (obj, reinterpret_cast <PyTypeObject*>(g_enum_Enum_type)))
919926 return ConstantKind::Enum;
920927
@@ -2097,6 +2104,7 @@ static void extract_identity_constant(PyObject* object, Vec<int64_t>* constants,
20972104 identity_constants->push_back (object);
20982105}
20992106
2107+
21002108static PyPtr parse_identity_constant_constraint (ConstantCursor& cursor,
21012109 const Vec<PyObject*>& identity_constants) {
21022110 int64_t address = cursor.next ();
@@ -2107,6 +2115,21 @@ static PyPtr parse_identity_constant_constraint(ConstantCursor& cursor,
21072115 CHECK_UNREACHABLE ;
21082116}
21092117
2118+ static Status extract_string_constant (PyObject* pyobj, Vec<int64_t >* constants,
2119+ Vec<PyObject*>* identity_constants,
2120+ Vec<PyPtr>* pyarg_refs) {
2121+ if (!PyUnicode_CHECK_INTERNED (pyobj)) {
2122+ PyPtr ref = newref (pyobj);
2123+ pyunicode_intern_in_place (&ref);
2124+ if (!PyUnicode_CHECK_INTERNED (ref.get ()))
2125+ return raise (PyExc_RuntimeError, " Failed to intern a string kernel argument" );
2126+ pyobj = ref.get ();
2127+ pyarg_refs->push_back (std::move (ref));
2128+ }
2129+ extract_identity_constant (pyobj, constants, identity_constants);
2130+ return OK ;
2131+ }
2132+
21102133static Status extract_foreign_dtype_constant (PyObject* object, Vec<int64_t >* constants,
21112134 Vec<PyObject*>* identity_constants) {
21122135 HashMap<PyPtr, ForeignDTypeInfo>::Item* item = get_foreign_dtype_registry ()->find (object);
@@ -2408,6 +2431,9 @@ static Status extract_arg(const DriverApi* driver, PyObject* obj, PythonArgKind
24082431 return OK ;
24092432 case PythonArgKind::ConstantNone:
24102433 return OK ;
2434+ case PythonArgKind::ConstantString:
2435+ return extract_string_constant (obj, &helper.constants , &helper.identity_constants ,
2436+ &helper.pyarg_refs );
24112437 case PythonArgKind::IdentityConstant:
24122438 extract_identity_constant (obj, &helper.constants , &helper.identity_constants );
24132439 return OK ;
0 commit comments