MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / __call__

Method __call__

dnn/scripts/gen_param_defs.py:385–477  ·  view source on GitHub ↗
(self, fout, defs)

Source from the content-addressed store, hash-verified

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"

Callers 3

__call__Method · 0.45
__call__Method · 0.45
__call__Method · 0.45

Calls 4

_writeMethod · 0.80
_get_headerMethod · 0.80
_processMethod · 0.80
joinMethod · 0.80

Tested by

no test coverage detected