| 383 | self._imperative = for_imperative |
| 384 | |
| 385 | def __call__(self, fout, defs): |
| 386 | super().__call__(fout) |
| 387 | self._enum_member2num = [] |
| 388 | self._write("# %s", self._get_header()) |
| 389 | self._write("import struct") |
| 390 | self._write("from . import enum36 as enum") |
| 391 | self._write( |
| 392 | "class _ParamDefBase:\n" |
| 393 | " def serialize(self):\n" |
| 394 | ' tag = struct.pack("I", type(self).TAG)\n' |
| 395 | " pdata = [getattr(self, i) for i in self.__slots__]\n" |
| 396 | " for idx, v in enumerate(pdata):\n" |
| 397 | " if isinstance(v, _EnumBase):\n" |
| 398 | " pdata[idx] = _enum_member2num[id(v)]\n" |
| 399 | " elif isinstance(v, _BitCombinedEnumBase):\n" |
| 400 | " pdata[idx] = v._value_\n" |
| 401 | " return tag + self._packer.pack(*pdata)\n" |
| 402 | "\n" |
| 403 | ) |
| 404 | # it's hard to mix custom implemention into enum, just do copy-paste instead |
| 405 | classbody = ( |
| 406 | " @classmethod\n" |
| 407 | " def __normalize(cls, val):\n" |
| 408 | " if isinstance(val, str):\n" |
| 409 | ' if not hasattr(cls, "__member_upper_dict__"):\n' |
| 410 | " cls.__member_upper_dict__ = {k.upper(): v\n" |
| 411 | " for k, v in cls.__members__.items()}\n" |
| 412 | " val = cls.__member_upper_dict__.get(val.upper(),val)\n" |
| 413 | " return val\n" |
| 414 | " @classmethod\n" |
| 415 | " def convert(cls, val):\n" |
| 416 | " val = cls.__normalize(val)\n" |
| 417 | " if isinstance(val, cls):\n" |
| 418 | " return val\n" |
| 419 | " return cls(val)\n" |
| 420 | " @classmethod\n" |
| 421 | " def _missing_(cls, value):\n" |
| 422 | " vnorm = cls.__normalize(value)\n" |
| 423 | " if vnorm is not value:\n" |
| 424 | " return cls(vnorm)\n" |
| 425 | " return super()._missing_(value)\n" |
| 426 | "\n" |
| 427 | ) |
| 428 | self._write("class _EnumBase(enum.Enum):\n" + classbody) |
| 429 | self._write("class _BitCombinedEnumBase(enum.Flag):\n" + classbody) |
| 430 | if not self._imperative: |
| 431 | self._write( |
| 432 | "def _as_dtype_num(dtype):\n" |
| 433 | " import megbrain.mgb as m\n" |
| 434 | " return m._get_dtype_num(dtype)\n" |
| 435 | "\n" |
| 436 | ) |
| 437 | |
| 438 | self._write( |
| 439 | "def _as_serialized_dtype(dtype):\n" |
| 440 | " import megbrain.mgb as m\n" |
| 441 | " return m._get_serialized_dtype(dtype)\n" |
| 442 | "\n" |