raahemnabeel commited on
Commit
bc53c10
·
verified ·
1 Parent(s): 7a28c91

Add files using upload-large-folder tool

Browse files
code/models/autoports/qwen_qwen3_5_122b_a10b/tt/model.py CHANGED
@@ -33,6 +33,7 @@ from models.autoports.qwen_qwen3_5_122b_a10b.tt.multichip_decoder import (
33
  MultichipDecoder,
34
  )
35
  from models.autoports.qwen_qwen3_5_122b_a10b.tt.optimized_decoder import Qwen35MoeOptimizedDecoderLayer
 
36
  from models.demos.blackhole.qwen36.tt.generator_interface import unpack_rope
37
  from models.demos.blackhole.qwen36.tt.model import Qwen36Model
38
  from models.tt_transformers.tt.common import Mode
@@ -61,6 +62,74 @@ class _CacheRoot(type(Path())):
61
  return super().__truediv__(key)
62
 
63
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
64
  class Qwen35MoeModel(Qwen36Model):
65
  """qwen36 full model with the optimized multichip MoE layer; see the module docstring for the decode contract."""
66
 
@@ -81,6 +150,14 @@ class Qwen35MoeModel(Qwen36Model):
81
  if not isinstance(tensor_cache_path, _CacheRoot):
82
  tensor_cache_path = _CacheRoot(tensor_cache_path or ".", disabled=tensor_cache_path is None)
83
  super().__init__(mesh_device, args, state_dict, tensor_cache_path)
 
 
 
 
 
 
 
 
84
  # Device-owned decode inputs: tokens via tt_out_tok, positions via plus_one (no per-token host refresh).
85
  self._tt_vllm_always_refresh_decode_trace_inputs = False
86
  rd = self.args.rope_head_dim
 
33
  MultichipDecoder,
34
  )
35
  from models.autoports.qwen_qwen3_5_122b_a10b.tt.optimized_decoder import Qwen35MoeOptimizedDecoderLayer
36
+ from models.common.sampling.generator import SeedManager
37
  from models.demos.blackhole.qwen36.tt.generator_interface import unpack_rope
38
  from models.demos.blackhole.qwen36.tt.model import Qwen36Model
39
  from models.tt_transformers.tt.common import Mode
 
62
  return super().__truediv__(key)
63
 
64
 
65
+ class _LockstepSeedManager(SeedManager):
66
+ """SeedManager whose UNSEEDED path pushes fresh per-slot seeds before EVERY draw.
67
+
68
+ The stock manager pushes concrete seeds once and then MAX_UINT32 ("SKIP"), which leaves each die's hardware PRNG
69
+ free-running. ``ttnn.sampling`` runs on EVERY die of the mesh (each die holds the same gathered candidate row), the
70
+ host reads die 0's token (``process_output_decode``) and every die feeds ITS OWN token back into the next decode
71
+ step (``tt_out_tok`` is the persistent decode token input). The free-running PRNGs stay in lockstep only while every
72
+ die advances its PRNG identically between draws. Inside a full decode step they can drift apart and draw DIFFERENT
73
+ tokens from the same distribution: measured on the sibling Qwen3-Coder-Next port (same sampler, same state) on
74
+ p300x2 at 20K context, 196 of 400 sampled steps had at least one die disagree with die 0. Die 0's token then goes
75
+ to the user while the model continues from the other dies' tokens (word corruption such as "checkpoints" +
76
+ "etensors"). See models/autoports/qwen_qwen3_coder_next/doc/text_quality/README.md.
77
+
78
+ Identical seeds pushed right before each draw (``ttnn.manual_seed`` runs immediately before ``ttnn.sampling`` in the
79
+ sampler trace, the seed tensor is replicated on the mesh) re-synchronise every die on every step: 0 % cross-die
80
+ disagreement and ~88 % step-to-step draw variation in the standalone 4-die op test. The seeded path is unchanged
81
+ (it already pushes concrete per-step seeds for every live slot). Cost: one 128-byte host-to-device copy per step.
82
+ """
83
+
84
+ def get_new_values(self, empty_slots=None, replicate_seeds=False):
85
+ if self._seed_active:
86
+ return super().get_new_values(empty_slots, replicate_seeds=replicate_seeds)
87
+ self._active_request_seed = False
88
+ new_seeds = [self._next_unseeded_device_seed() for _ in range(self.max_batch_size)]
89
+ self.write_device_seed_values(new_seeds)
90
+ self._reseted = False
91
+ self._needs_skip = False
92
+ return tuple(new_seeds)
93
+
94
+
95
+ def _install_lockstep_sampling(sampling, mesh_device):
96
+ """Make every die of the mesh feed back the SAME sampled token the host reads (see _LockstepSeedManager).
97
+
98
+ Two independent guarantees: (1) identical fresh seeds before every draw, so the dies draw identically; (2) after
99
+ ``ttnn.sampling`` the token row is all-gathered and die 0's row is copied into the token buffer on every die, so
100
+ even a die whose draw differed for any other reason feeds back die 0's token - the one the host reports.
101
+ QWEN35_SAMPLER_TOKEN_BROADCAST=0 disables (2) (A/B only). The greedy force-argmax path (not enabled on this mesh)
102
+ is deterministic per die and is left alone.
103
+ """
104
+ old = sampling.seed_manager
105
+ sampling.seed_manager = _LockstepSeedManager(
106
+ sampling.tt_sampling, max_batch_size=old.max_batch_size, salt_duplicate_seeds=old.salt_duplicate_seeds
107
+ )
108
+ ts = sampling.tt_sampling
109
+ orig_forward = ts.forward
110
+ broadcast = os.environ.get("QWEN35_SAMPLER_TOKEN_BROADCAST", "1") == "1"
111
+
112
+ def forward(x, tt_out_tok=None):
113
+ tokens, log_probs = orig_forward(x, tt_out_tok=tt_out_tok)
114
+ if not broadcast or ts.force_argmax_sampling:
115
+ return tokens, log_probs
116
+ # [1,1,1,S] uint32 ROW_MAJOR on every die -> [1,1,N,S] -> die 0's row -> back into the (persistent) buffer.
117
+ # (The sampled-token log-probs, when enabled, were computed from each die's own draw; on this mesh the plugin
118
+ # samples on the host whenever log-probs are requested, so that tensor is never consumed here.)
119
+ gathered = ttnn.all_gather(tokens, dim=2, memory_config=ttnn.DRAM_MEMORY_CONFIG)
120
+ first = ttnn.slice(gathered, [0, 0, 0, 0], [1, 1, 1, tokens.shape[-1]])
121
+ ttnn.deallocate(gathered)
122
+ ttnn.copy(first, tokens)
123
+ ttnn.deallocate(first)
124
+ return tokens, log_probs
125
+
126
+ ts.forward = forward
127
+ logger.info(
128
+ f"sampler lockstep installed on {mesh_device.get_num_devices()} dies: fresh identical seeds every draw"
129
+ f"{', die-0 token broadcast' if broadcast else ''}"
130
+ )
131
+
132
+
133
  class Qwen35MoeModel(Qwen36Model):
134
  """qwen36 full model with the optimized multichip MoE layer; see the module docstring for the decode contract."""
135
 
 
150
  if not isinstance(tensor_cache_path, _CacheRoot):
151
  tensor_cache_path = _CacheRoot(tensor_cache_path or ".", disabled=tensor_cache_path is None)
152
  super().__init__(mesh_device, args, state_dict, tensor_cache_path)
153
+ if (
154
+ self.sampling is not None
155
+ and mesh_device.get_num_devices() > 1
156
+ and os.environ.get("QWEN35_SAMPLER_LOCKSTEP", "1") == "1"
157
+ ):
158
+ # every die must feed back the token the host reads (the Qwen3-Coder-Next sibling port's sampled-output
159
+ # word corruption, same sampler and seed state here); QWEN35_SAMPLER_LOCKSTEP=0 restores the stock sampler
160
+ _install_lockstep_sampling(self.sampling, mesh_device)
161
  # Device-owned decode inputs: tokens via tt_out_tok, positions via plus_one (no per-token host refresh).
162
  self._tt_vllm_always_refresh_decode_trace_inputs = False
163
  rd = self.args.rope_head_dim