| 15 | |
| 16 | |
| 17 | def test_metadata(): |
| 18 | x = Tensor(0) |
| 19 | |
| 20 | @trace(symbolic=True, capture_as_const=True) |
| 21 | def fwd(x): |
| 22 | return x * 2 |
| 23 | |
| 24 | fwd(x) |
| 25 | |
| 26 | orig_model = io.BytesIO() |
| 27 | fwd.dump(orig_model, user_info="test", optimize_for_inference=False) |
| 28 | orig_model.seek(0) |
| 29 | graph = Net.load(orig_model) |
| 30 | assert graph.metadata == { |
| 31 | "user_info": "test", |
| 32 | "graph_modified": False, # False: tracing.dump |
| 33 | "optimized_for_inference": False, |
| 34 | } |
| 35 | |
| 36 | orig_model.seek(0) |
| 37 | graph.dump( |
| 38 | orig_model, |
| 39 | user_info={"str": "x", "tensor": x, "module": M.Module, "none": None}, |
| 40 | optimize_for_inference=True, |
| 41 | enable_nchw4=True, |
| 42 | enable_ioc16=True, |
| 43 | ) |
| 44 | orig_model.seek(0) |
| 45 | graph = Net.load(orig_model) |
| 46 | assert graph.metadata == { |
| 47 | "user_info": {"str": "x", "tensor": x, "module": M.Module, "none": None}, |
| 48 | "graph_modified": True, # True: Network.dump |
| 49 | "optimized_for_inference": True, |
| 50 | "enable_fuse_grain": True, |
| 51 | "enable_nchw4": True, |
| 52 | "enable_ioc16": True, |
| 53 | } |
| 54 | |
| 55 | orig_model.seek(0) |
| 56 | fwd.dump(orig_model, enable_metadata=False) |
| 57 | orig_model.seek(0) |
| 58 | graph = Net.load(orig_model) |
| 59 | assert graph.metadata is None |
| 60 | |
| 61 | |
| 62 | def test_replace_var(): |